Skip to content

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

Open
parsley9877 wants to merge 1 commit into
mainfrom
WeaverOmniTransformer-bring-up
Open

parsley9877 wants to merge 1 commit into
mainfrom
WeaverOmniTransformer-bring-up

Conversation

@parsley9877

@parsley9877 parsley9877 commented Sep 30, 2026 •

Copy link
Copy Markdown
Collaborator

Summary# 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

All unit tests, golden numerical parity checks, and 36-layer TPU v5p execution tests in tests/unit/weaver_layers_test.py pass:

Test Suite / Case Device Precision Status Notes / Numerical Parity
WeaverOmniTransformerUnitTest::test_patchify_unpatchify_roundtrip_and_shapes CPU / TPU float32 ✅ Passed Exact round-trip reconstruction (atol=1e-6) for 2D & 3D patch tensors
WeaverOmniTransformerUnitTest::test_timestep_embedder_shapes_and_float32_sinusoidal CPU / TPU bfloat16 / float32 ✅ Passed Verifies float32 sinusoidal embedding and [B] / [B, T] timestep broadcasting
WeaverOmniTransformerUnitTest::test_build_weaver_3d_position_ids CPU / TPU int32 ✅ Passed Exact match on packed text (i, i, i) and 3D vision (S_und + t, h, w) coordinates
WeaverOmniTransformerUnitTest::test_weaver_mini_diffuser_config_loading CPU / TPU N/A ✅ Passed Verifies weaver-mini-diffuser.yml + isolated weaver.yml populate WeaverConfig
WeaverOmniTransformerUnitTest::test_scan_layers_false_and_true_equivalence CPU / TPU float32 ✅ Passed scan_layers=False vs scan_layers=True match within rtol=1e-5, atol=1e-5
WeaverOmniTransformerGoldenParityTest::test_e2e_golden_parity (weaver_mini, scan_layers=False) CPU bfloat16 ✅ Passed preds_avg_abs=3.26e-03, hidden_avg_abs=6.91e-03 vs GPU PyTorch golden
WeaverOmniTransformerGoldenParityTest::test_e2e_golden_parity (weaver_mini, scan_layers=True) CPU bfloat16 ✅ Passed preds_avg_abs=3.26e-03, hidden_avg_abs=6.91e-03 vs GPU PyTorch golden
WeaverOmniTransformerGoldenParityTest::test_e2e_golden_parity (weaver_max, scan_layers=False) CPU bfloat16 ✅ Passed preds_avg_abs=9.98e-05, hidden_avg_abs=3.84e-03 vs GPU PyTorch golden
WeaverOmniTransformerGoldenParityTest::test_e2e_golden_parity (weaver_max, scan_layers=True) CPU bfloat16 ✅ Passed preds_avg_abs=9.98e-05, hidden_avg_abs=3.84e-03 vs GPU PyTorch golden
WeaverOmniTransformerTPUTest::test_full_36_layer_backbone_forward_tpu (unscanned & scanned) TPU v5p-8 bfloat16 ✅ Passed Full 36-layer (hidden_size=4096) @jax.jit forward pass ([2, 48, 2, 4, 4], all finite)
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 23.52941% with 195 lines in your changes missing coverage. Please review.

Files with missing lines Patch % Lines
src/maxtext/models/weaver.py 23.22% 194 Missing and 1 partial ⚠️

📢 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.

Comment thread tests/unit/weaver_transformer_test.py Outdated

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.

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

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

Done. merged all WeaverOmniTransformer unit, golden parity, and TPU tests into tests/unit/weaver_layers_test.py, marked them @pytest.mark.scheduled_only, and removed tests/unit/weaver_transformer_test.py.

Comment thread tests/unit/weaver_transformer_test.py Outdated
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?

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

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

Yes, the original PyTorch reference requires a GPU/CUDA environment (e.g., FlashAttention/CUDA kernels) unavailable on CPU/TPU CI runners, and (2) instantiating and running the PyTorch model in-process slows down the test suite compared to loading a small .npz archive.

# 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"?

Comment thread src/maxtext/configs/types.py Outdated
"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

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

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

Done. renamed weaver-nano-diffuser to weaver-mini-diffuser (src/maxtext/configs/models/weaver-mini-diffuser.yml) and updated types.py and the tests accordingly.

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.

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

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

Done. simplified the override handling in WeaverOmniTransformer.init using dataclasses.fields(WeaverConfig) to avoid redundant field-by-field conversion.

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! It is still a bit confusing on from_config and from_matext_config.
e.g. head_dim=int(getattr(config, "head_dim", 128)),
But you can continue using this as long as it is in the self-contained model file.
Besides, it would be appreciated if you can add clarification function comments for them.

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

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

Thank you! Added comments.

Comment thread src/maxtext/models/models.py Outdated
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?

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

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

Thank you!. that was only a convenience re-export and isn't needed in models.py. Removed the import and alias so models.py is untouched.

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?

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

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

Done. moved latent_channels, patch_size, time_embed_in_channels, timestep_scale, timestep_max_period, qk_norm_for_text, qk_norm_for_diffusion, use_und_k_norm_for_gen, weaver_block_q, and weaver_block_kv into base.yml, types.py, and weaver-mini-diffuser.yml, and wired them through WeaverConfig, attention kernels, and WeaverTimeEmbedder.

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 replies!
Some configs are added in core files including base.yml and types.py. Can we please isolate them by adding in weaver.yml?

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

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

Sure.

-Reverted src/maxtext/configs/base.yml completely (0 changes to base.yml).
-Removed all 10 Weaver-specific diffusion and pathway fields (latent_channels, patch_size, time_embed_in_channels, timestep_scale, timestep_max_period, qk_norm_for_text, qk_norm_for_diffusion, use_und_k_norm_for_gen, weaver_block_q, weaver_block_kv) from MultimodalGeneral in src/maxtext/configs/types.py.
-Isolated those Weaver-specific parameters into src/maxtext/configs/models/weaver.yml, which WeaverConfig.from_maxtext_config loads directly via _load_weaver_yaml_defaults() in src/maxtext/models/weaver.py.

parsley9877 added a commit that referenced this pull request Oct 5, 2026
…arams to config, merge tests into weaver_layers_test.py
@parsley9877

Copy link
Copy Markdown
Collaborator Author

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?

Shared the reproduction setup with you offline.

@parsley9877

Copy link
Copy Markdown
Collaborator Author

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

Thank you for your comment. Checked the Gemini review comments (gemini-code-assist had no inline comments on this PR, and I also addressed the remaining Gemini suggestion from #5384 on _adapt_kernel_init using inspect.signature).
Marked all the WeaverOmniTransformer test classes (WeaverOmniTransformerUnitTest, WeaverOmniTransformerGoldenParityTest, and WeaverOmniTransformerTPUTest) with @pytest.mark.scheduled_only in tests/unit/weaver_layers_test.py.

@parsley9877
parsley9877 force-pushed the WeaverOmniTransformer-bring-up branch from f2ea885 to da02a6d Compare October 5, 2026 16:38
@lydhr

lydhr commented Oct 6, 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?

Shared the reproduction setup with you offline.

Thanks! You can feel free to add summary of key results and elaboration on testing scripts here, referring to the previous PRs.

@parsley9877
parsley9877 force-pushed the WeaverOmniTransformer-bring-up branch from da02a6d to 04ebc62 Compare October 6, 2026 17:16
@parsley9877

Copy link
Copy Markdown
Collaborator Author

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?

Shared the reproduction setup with you offline.

Thanks! You can feel free to add summary of key results and elaboration on testing scripts here, referring to the previous PRs.

Updated the PR description to include test results.

@parsley9877
parsley9877 force-pushed the WeaverOmniTransformer-bring-up branch from 04ebc62 to 587e102 Compare October 6, 2026 17:27

@lydhr lydhr 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.

LGTM. Thanks for your replies!

Does CI failures make sense? If not, please leave a comment with some clarification and analysis and we can manually skip them.

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