Add WeaverOmniTransformer full 36-layer backbone architecture and tests - #5476
parsley9877 wants to merge 2 commits 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.
| 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?
| # 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
| 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.
| 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?
| 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?
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
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.