Skip to content

Add WeaverOmniTransformer full 36-layer backbone architecture and tests - #5476

Open
parsley9877 wants to merge 2 commits into
mainfrom
WeaverOmniTransformer-bring-up
Open

parsley9877 wants to merge 2 commits into
mainfrom
WeaverOmniTransformer-bring-up

Conversation

@parsley9877

@parsley9877 parsley9877 commented Sep 30, 2026 •

Copy link
Copy Markdown
Collaborator

Description

Assembles the full 36-layer WeaverOmniTransformer Mixture-of-Transformers (MoT) diffusion backbone in Flax NNX on top of the existing WeaverMoTDecoderLayer and WeaverJointAttention modules.

  • src/maxtext/models/weaver.py: Implements WeaverOmniTransformer, WeaverTimeEmbedder, patchify_latents/unpatchify_latents, and build_weaver_3d_position_ids with support for both scan_layers=False and scan_layers=True.
  • src/maxtext/configs/models/weaver-nano-diffuser.yml: Adds the 36-layer diffuser model config and registers weaver / weaver-nano-diffuser in common_types.py, types.py, nnx_decoders.py, and models.py.
  • tests/unit/weaver_transformer_test.py: Adds CPU unit tests, E2E golden numerical parity tests, and scheduled TPU forward-pass tests.

Tests

pytest tests/unit/weaver_transformer_test.py -k "not TPUTest"

Checklist

Before submitting this PR, please make sure (put X in square brackets):

  • I have performed a self-review of my code. For an optional AI review, add the gemini-review label.
  • I have necessary comments in my code, particularly in hard-to-understand areas.
  • I have run end-to-end tests tests and provided workload links above if applicable.
  • I have made or will make corresponding changes to the doc if needed, including adding new documentation pages to the relevant Table of Contents (toctree directive) as explained in our documentation.

@gemini-code-assist gemini-code-assist Bot left a comment

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Code Review

This pull request implements the Weaver Omni-Transformer (WeaverOmniTransformer) architecture, a 36-layer joint multimodal diffusion backbone. Key additions include the WEAVER decoder block type, the weaver-nano-diffuser model configuration, latent patchification/unpatchification helpers, timestep embedding, and 3D M-RoPE position ID generation. Unit, golden end-to-end parity, and TPU v5 tests are also added to validate the implementation. No review comments were provided, and there is no feedback to address.

@codecov

codecov Bot commented Sep 30, 2026 •

Copy link
Copy Markdown

Codecov Report

❌ Patch coverage is 78.16594% with 50 lines in your changes missing coverage. Please review.

Files with missing lines Patch % Lines
src/maxtext/models/weaver.py 77.87% 27 Missing and 23 partials ⚠️

📢 Thoughts on this report? Let us know!

@parsley9877
parsley9877 added this pull request to stack #5489 October 1, 2026 17:24
@lydhr

lydhr commented Oct 1, 2026

Copy link
Copy Markdown
Collaborator

QQ: did you run the tests/unit/weaver_layers_test.py as the previous PR?
Can u please elaborate on the setup&steps of your GPU golden parity checks?

@hengtaoguo hengtaoguo left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Thanks for the work! Could you also check the gemini review comments? We should also mark these tests schedule-only.

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I wonder if you could merge these tests into weaver_layers_test.py? So that every model only corresponds to one test. Or let me know if you prefer to isolate it.

self.skipTest(
f"Golden test asset {_TRANSFORMER_GOLDEN_FILENAME} not found locally under /tmp "
"and could not be downloaded from gs://maxtext-test-assets/."
)

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Just to confirm, is the golden reference data pre-dumped in this GCS bucket? Also, if we import the original PyTorch module here, will that significantly slow down the test suite?

# Model config for Weaver Nano Diffuser (36-layer Mixture-of-Transformers Joint Backbone)

# Core Architectural Parameters
decoder_block: "weaver"

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Should we use "decoder_block" or a new "diffuser_block"?

"maxtext-omni-gemma3-qwen3",
"weaver-mini",
"weaver-max",
"weaver-nano-diffuser",

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

maybe weaver-mini-diffuser to align with above? We should also update the config file name correspondingly

self.weight_dtype = _resolve_jax_dtype(self.weight_dtype)

@classmethod
def from_maxtext_config(cls, config: Any) -> WeaverConfig:

@lydhr lydhr Oct 3, 2026 •

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This config conversion seems having redundancy, but you can continue using this as long as it is in the self-contained model file.
Ditto for the config conversion below.

from maxtext.layers.encoders import AudioEncoder, VisionEncoder
from maxtext.layers.multi_token_prediction import MultiTokenPredictionBlock
from maxtext.layers.quantizations import AqtQuantization as Quant
from maxtext.models import weaver

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Before this PR, models.py imported nothing from maxtext.models. Can u please confirm if the new import exists only to create a re-export?

base_num_decoder_layers: 36
head_dim: 128
mlp_activations: ["silu", "linear"]
vocab_size: 151936

@lydhr lydhr Oct 3, 2026 •

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

There are many hardcoded length of latent_channels, token, timestep_scale, etc., (some are in the previous checked-in PR) in weaver.py.

Can you please clean them up as much as possible by putting them in the config?

This branch has not been deployed

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

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

3 participants