Skip to content

feat(nnx): Missing object-oriented pooling layers in NNX #5202

Description

@divye-joshi

Problem Description

Currently, flax.nnx lacks native object-oriented pooling modules such as MaxPool, AvgPool, and GlobalAveragePool. Users migrating from frameworks like Keras or PyTorch—or even transitioning from flax.linen—are forced to mix functional API calls within the object-oriented NNX structure. This creates an inconsistent developer experience and requires manual boilerplate for common operations like Global Average Pooling.

Proposed Feature

Introduce a dedicated pooling module suite within nnx that mirrors the ergonomic design of other NNX layers. This includes:

  • Subsampling Modules: MaxPool, AvgPool, and MinPool.
  • Global Pooling: A dedicated GlobalAveragePool module to replace manual jnp.mean calls.

Implementation Status

I have already implemented these modules and exposed them in the nnx namespace.
See Pull Request: #5201

Justification & Benefits

  1. API Consistency: Maintains the OO-flow of NNX without jumping back into linen.functional.
  2. Framework Parity: Lowers the barrier for users migrating from Keras/PyTorch.
  3. Readability: Simplifies model definitions, especially for standard CNN architectures.

Activity

  1. changed the title [-]feat(nnx) : Missing object-oriented pooling layers in NNX[/-] [+]feat(nnx): Missing object-oriented pooling layers in NNX[/+] on Jan 26, 2026
  2. vfdev-5 commented on Jan 26, 2026

    @vfdev-5
    Collaborator

    @starryendymion thanks for the report and the PR. Previously we had a PR exposing pooling ops in nnx: #5057 but unfortunately it was reverted as it introduced some internal errors. We may want to first reintroduce what it was doing in a safe manner (without changing linen codebase).

    As for global avg pool we can do:

    gpool = lambda x: nnx.avg_pool(x, (x.shape[1], x.shape[2])). #  x is in NHWC format

    But I agree that object-oriented pooling modules such as MaxPool, AvgPool, and GlobalAveragePool could be good to have, especially for users migrating from frameworks like Keras or PyTorch.

  3. divye-joshi commented on Jan 26, 2026

    @divye-joshi
    Author

    @vfdev-5

    Thanks for the context regarding PR #5057!

    To ensure safety, this PR is purely additive. It creates a new flax/nnx/nn/pooling.py module that wraps the existing, stable flax.linen.pooling functional API.

    It does not modify any code within flax/linen, so it should avoid the regression issues encountered previously.

    Regarding GlobalAveragePool, I agree the lambda works, but having it as a dedicated layer makes migration from Keras/PyTorch much smoother (one less shape calculation for the user to worry about).

    I've verified the shapes and values match the Linen functional equivalents locally. Please let me know if there are any specific tests you'd like me to add!

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

No one assigned

    Labels

    No labels
    No labels

    Type

    No type

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions