Repository navigation
Add WeaverOmniTransformer full 36-layer backbone architecture and tests - #5476
parsley9877 wants to merge 1 commit into
Conversation
There was a problem hiding this comment.
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 Report❌ Patch coverage is
📢 Thoughts on this report? Let us know! |
|
QQ: did you run the tests/unit/weaver_layers_test.py as the previous PR? |
hengtaoguo
left a comment
There was a problem hiding this comment.
Thanks for the work! Could you also check the gemini review comments? We should also mark these tests schedule-only.
There was a problem hiding this comment.
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.
There was a problem hiding this comment.
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.
| self.skipTest( | ||
| f"Golden test asset {_TRANSFORMER_GOLDEN_FILENAME} not found locally under /tmp " | ||
| "and could not be downloaded from gs://maxtext-test-assets/." | ||
| ) |
There was a problem hiding this comment.
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?
There was a problem hiding this comment.
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" |
There was a problem hiding this comment.
Should we use "decoder_block" or a new "diffuser_block"?
| "maxtext-omni-gemma3-qwen3", | ||
| "weaver-mini", | ||
| "weaver-max", | ||
| "weaver-nano-diffuser", |
There was a problem hiding this comment.
maybe weaver-mini-diffuser to align with above? We should also update the config file name correspondingly
There was a problem hiding this comment.
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: |
There was a problem hiding this comment.
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.
There was a problem hiding this comment.
Done. simplified the override handling in WeaverOmniTransformer.init using dataclasses.fields(WeaverConfig) to avoid redundant field-by-field conversion.
There was a problem hiding this comment.
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.
There was a problem hiding this comment.
Thank you! Added comments.
| 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 |
There was a problem hiding this comment.
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?
There was a problem hiding this comment.
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 |
There was a problem hiding this comment.
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?
There was a problem hiding this comment.
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.
There was a problem hiding this comment.
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?
There was a problem hiding this comment.
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.
…arams to config, merge tests into weaver_layers_test.py
Shared the reproduction setup with you offline. |
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). |
f2ea885 to
da02a6d
Compare
Thanks! You can feel free to add summary of key results and elaboration on testing scripts here, referring to the previous PRs. |
da02a6d to
04ebc62
Compare
Updated the PR description to include test results. |
04ebc62 to
587e102
Compare
Summary# Description
Assembles the full 36-layer
WeaverOmniTransformerMixture-of-Transformers (MoT) diffusion backbone in Flax NNX on top of the existingWeaverMoTDecoderLayerandWeaverJointAttentionmodules.src/maxtext/models/weaver.py: ImplementsWeaverOmniTransformer,WeaverTimeEmbedder,patchify_latents/unpatchify_latents, andbuild_weaver_3d_position_idswith support for bothscan_layers=Falseandscan_layers=True.src/maxtext/configs/models/weaver-nano-diffuser.yml: Adds the 36-layer diffuser model config and registersweaver/weaver-nano-diffuserincommon_types.py,types.py,nnx_decoders.py, andmodels.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.pypass:WeaverOmniTransformerUnitTest::test_patchify_unpatchify_roundtrip_and_shapesfloat32atol=1e-6) for 2D & 3D patch tensorsWeaverOmniTransformerUnitTest::test_timestep_embedder_shapes_and_float32_sinusoidalbfloat16/float32float32sinusoidal embedding and[B]/[B, T]timestep broadcastingWeaverOmniTransformerUnitTest::test_build_weaver_3d_position_idsint32(i, i, i)and 3D vision(S_und + t, h, w)coordinatesWeaverOmniTransformerUnitTest::test_weaver_mini_diffuser_config_loadingweaver-mini-diffuser.yml+ isolatedweaver.ymlpopulateWeaverConfigWeaverOmniTransformerUnitTest::test_scan_layers_false_and_true_equivalencefloat32scan_layers=Falsevsscan_layers=Truematch withinrtol=1e-5, atol=1e-5WeaverOmniTransformerGoldenParityTest::test_e2e_golden_parity (weaver_mini, scan_layers=False)bfloat16preds_avg_abs=3.26e-03,hidden_avg_abs=6.91e-03vs GPU PyTorch goldenWeaverOmniTransformerGoldenParityTest::test_e2e_golden_parity (weaver_mini, scan_layers=True)bfloat16preds_avg_abs=3.26e-03,hidden_avg_abs=6.91e-03vs GPU PyTorch goldenWeaverOmniTransformerGoldenParityTest::test_e2e_golden_parity (weaver_max, scan_layers=False)bfloat16preds_avg_abs=9.98e-05,hidden_avg_abs=3.84e-03vs GPU PyTorch goldenWeaverOmniTransformerGoldenParityTest::test_e2e_golden_parity (weaver_max, scan_layers=True)bfloat16preds_avg_abs=9.98e-05,hidden_avg_abs=3.84e-03vs GPU PyTorch goldenWeaverOmniTransformerTPUTest::test_full_36_layer_backbone_forward_tpu (unscanned & scanned)bfloat16hidden_size=4096)@jax.jitforward 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):
gemini-reviewlabel.