Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions src/maxtext/common/common_types.py
Original file line number Diff line number Diff line change
Expand Up @@ -141,6 +141,7 @@ class DecoderBlockType(enum.Enum):
OLMO3 = "olmo3"
DEEPSEEK4 = "deepseek4"
ENVY = "envy"
WEAVER = "weaver"


class VisionEncoderBlockType(enum.Enum):
Expand Down
46 changes: 46 additions & 0 deletions src/maxtext/configs/models/weaver-mini-diffuser.yml
Original file line number Diff line number Diff line change
@@ -0,0 +1,46 @@
# Copyright 2023–2026 Google LLC
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# https://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.

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

# Core Architectural Parameters
decoder_block: "weaver"
Comment thread
parsley9877 marked this conversation as resolved.
base_emb_dim: 4096
base_mlp_dim: 12288
base_num_query_heads: 32
base_num_kv_heads: 8
base_num_decoder_layers: 36
head_dim: 128
mlp_activations: ["silu", "linear"]
vocab_size: 151936
Comment thread
parsley9877 marked this conversation as resolved.
normalization_layer_epsilon: 1.0e-6
use_qk_norm: true
logits_via_embedding: false

# RoPE Settings
rope_max_timescale: 5000000

# General Model Settings
enable_dropout: false

# Multimodal & 3D M-RoPE Settings
use_multimodal: true
use_mrope: true
mrope_section: [24, 20, 20]

# Note: Weaver-specific diffusion & pathway parameters (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)
# are isolated in src/maxtext/configs/models/weaver.yml and loaded by WeaverConfig.

50 changes: 50 additions & 0 deletions src/maxtext/configs/models/weaver.yml
Original file line number Diff line number Diff line change
@@ -0,0 +1,50 @@
# Copyright 2023–2026 Google LLC
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# https://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.

# Isolated Weaver Mixture-of-Transformers diffusion & pathway configuration.
# Loaded directly by WeaverConfig in src/maxtext/models/weaver.py so that
# Weaver-specific diffusion parameters stay isolated from base.yml and types.py.

# Core Architectural Parameters
decoder_block: "weaver"
base_emb_dim: 4096
base_mlp_dim: 12288
base_num_query_heads: 32
base_num_kv_heads: 8
base_num_decoder_layers: 36
head_dim: 128
mlp_activations: ["silu", "linear"]
vocab_size: 151936
normalization_layer_epsilon: 1.0e-6
use_qk_norm: true
logits_via_embedding: false

# RoPE & 3D M-RoPE Settings
rope_max_timescale: 5000000
enable_dropout: false
use_multimodal: true
use_mrope: true
mrope_section: [24, 20, 20]

# Weaver Diffusion & Dual-Pathway Settings
latent_channels: 48
patch_size: 2
time_embed_in_channels: 256
timestep_scale: 0.001
timestep_max_period: 10000
qk_norm_for_text: true
qk_norm_for_diffusion: true
use_und_k_norm_for_gen: false
weaver_block_q: 128
weaver_block_kv: 128
2 changes: 2 additions & 0 deletions src/maxtext/configs/types.py
Original file line number Diff line number Diff line change
Expand Up @@ -292,6 +292,7 @@ class ProfilerType(str, Enum):
"qwen3-vl-30b-a3b",
"weaver-mini",
"weaver-max",
"weaver-mini-diffuser",
"qwen3-next-80b-a3b",
"qwen3-omni-30b-a3b",
"qwen3-custom-30b-a3b",
Expand Down Expand Up @@ -5407,6 +5408,7 @@ def calculate_global_batch_sizes(per_device_batch_size, expansion_factor, num_de
"maxtext-omni-gemma3-qwen3",
"weaver-mini",
"weaver-max",
"weaver-mini-diffuser",
)
if self.model_name not in valid_mm_models and self.model_name != "default":
raise ValueError(f"Multimodal is only supported for {valid_mm_models}, not {self.model_name}")
Expand Down
2 changes: 2 additions & 0 deletions src/maxtext/layers/nnx_decoders.py
Original file line number Diff line number Diff line change
Expand Up @@ -69,6 +69,7 @@
qwen3_5,
qwen3_custom,
simple_layer,
weaver,
)

from maxtext.multimodal import utils as mm_utils
Expand Down Expand Up @@ -1240,6 +1241,7 @@ def get_deepseek():
DecoderBlockType.LLAMA4: get_scannable(llama4.Llama4DecoderLayer, llama4.Llama4ScannableBlock),
DecoderBlockType.OLMO3: get_scannable(olmo3.Olmo3DecoderLayer, olmo3.Olmo3ScannableBlock),
DecoderBlockType.ENVY: get_scannable(envy.EnvyDecoderLayer, envy.EnvyScannableBlock),
DecoderBlockType.WEAVER: [weaver.WeaverMoTDecoderLayer],
}

if cfg.decoder_block not in layer_map:
Expand Down
Loading
Loading