From 2b31681c289c2db6480c75171295f6a570869e6d Mon Sep 17 00:00:00 2001 From: haibo Date: Mon, 19 Jan 2026 12:56:45 +0800 Subject: [PATCH 01/18] Qwen2 support --- bonsai/models/qwen2/__init__.py | 0 bonsai/models/qwen2/modeling.py | 304 ++++++++++++++++++++++++++++++++ bonsai/models/qwen2/params.py | 146 +++++++++++++++ 3 files changed, 450 insertions(+) create mode 100644 bonsai/models/qwen2/__init__.py create mode 100644 bonsai/models/qwen2/modeling.py create mode 100644 bonsai/models/qwen2/params.py diff --git a/bonsai/models/qwen2/__init__.py b/bonsai/models/qwen2/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/bonsai/models/qwen2/modeling.py b/bonsai/models/qwen2/modeling.py new file mode 100644 index 00000000..7b5c71ff --- /dev/null +++ b/bonsai/models/qwen2/modeling.py @@ -0,0 +1,304 @@ +# Copyright 2025 The JAX Authors. +# +# 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. + +import dataclasses +import math + +import jax +from flax import nnx +from jax import P +from jax import numpy as jnp +from jax.sharding import PartitionSpec, get_abstract_mesh +from jaxtyping import Array + +from bonsai.models.qwen3.modeling import ( + Cache, + LayerCache, + MLP, + RMSNorm, + ShardingCfg, + ShardingSpec, + _generate_pos_embeddings, + apply_rope, + compute_positions_from_segment_ids, + count_left_pads, + count_right_pads, + reshard, + shard, +) + +_K_MASK = jnp.finfo(jnp.bfloat16).min + + +@dataclasses.dataclass(frozen=True) +class ModelConfig: + num_layers: int + vocab_size: int + emb_dim: int + mlp_dim: int + num_heads: int + head_dim: int + num_kv_heads: int + rope_theta: int + norm_eps: float + tie_word_embeddings: bool + use_causal_mask: bool = True # False for bidirectional attention + shd_cfg: ShardingCfg = ShardingCfg.no_sharding() + + @classmethod + def _from_param(cls, use_sharding: bool, **kwargs): + if use_sharding: + kwargs["shd_cfg"] = ShardingCfg.default() + return cls(**kwargs) + + @classmethod + def qwen2_0_5b(cls, use_sharding: bool = False): # qwen2.5-0.5B + return cls._from_param( + use_sharding, + num_layers=24, + vocab_size=151936, + emb_dim=896, + mlp_dim=4864, + num_heads=14, + head_dim=64, + num_kv_heads=2, + norm_eps=1e-06, + rope_theta=1000000, + tie_word_embeddings=True, + ) + + @classmethod + def qwen2_1_5b(cls, use_sharding: bool = False): # qwen2-1.5B + return cls._from_param( + use_sharding, + num_layers=28, + vocab_size=151936, + emb_dim=1536, + mlp_dim=8960, + num_heads=12, + head_dim=128, + num_kv_heads=2, + norm_eps=1e-06, + rope_theta=1000000, + tie_word_embeddings=True, + ) + + + @classmethod + def qwen2_7b(cls, use_sharding: bool = False): # qwen2-7B + return cls._from_param( + use_sharding, + num_layers=28, + vocab_size=152064, + emb_dim=3584, + mlp_dim=18944, + num_heads=28, + head_dim=128, + num_kv_heads=4, + norm_eps=1e-06, + rope_theta=1000000, + tie_word_embeddings=False, + ) + + @classmethod + def qwen2_72b(cls, use_sharding: bool = False): # qwen2-72B + return cls._from_param( + use_sharding, + num_layers=80, + vocab_size=151936, + emb_dim=8192, + mlp_dim=29568, + num_heads=64, + head_dim=128, + num_kv_heads=8, + norm_eps=1e-06, + rope_theta=1000000, + tie_word_embeddings=False, + ) + + +class Attention(nnx.Module): + def __init__(self, cfg: ModelConfig, *, rngs: nnx.Rngs): + self.shd_cfg = cfg.shd_cfg + + # Standard Linear layers matching official Qwen2 implementation + # q_proj: [B, T, D] @ [D, N*H] -> [B, T, N*H] + self.q_proj = shard( + nnx.Linear(cfg.emb_dim, cfg.num_heads * cfg.head_dim, use_bias=True, rngs=rngs), + self.shd_cfg.q_weight_ndh + ) + # k_proj: [B, T, D] @ [D, K*H] -> [B, T, K*H] + self.k_proj = shard( + nnx.Linear(cfg.emb_dim, cfg.num_kv_heads * cfg.head_dim, use_bias=True, rngs=rngs), + self.shd_cfg.kv_weight_ndh + ) + # v_proj: [B, T, D] @ [D, K*H] -> [B, T, K*H] + self.v_proj = shard( + nnx.Linear(cfg.emb_dim, cfg.num_kv_heads * cfg.head_dim, use_bias=True, rngs=rngs), + self.shd_cfg.kv_weight_ndh + ) + # o_proj: [B, T, N*H] @ [N*H, D] -> [B, T, D] + self.o_proj = shard( + nnx.Linear(cfg.num_heads * cfg.head_dim, cfg.emb_dim, use_bias=False, rngs=rngs), + self.shd_cfg.o_weight_nhd + ) + + self.cfg = cfg + self.n_rep = cfg.num_heads // cfg.num_kv_heads + self.scale = cfg.head_dim**-0.5 + self.rope_theta = cfg.rope_theta + + @jax.named_scope("attention") + def __call__(self, x: Array, cache: LayerCache | None, segment_ids: Array) -> Array: + # Linear projections output [B, T, N*H] or [B, T, K*H], then reshape to [B, T, N/K, H] + b, t = x.shape[:2] + + query_proj = self.q_proj(x).reshape(b, t, self.num_heads, self.head_dim) + query_proj = shard(query_proj, self.shd_cfg.act_btnh) # [B, T, N, H] + + key_proj = self.k_proj(x).reshape(b, t, self.num_kv_heads, self.head_dim) + key_proj = shard(key_proj, self.shd_cfg.act_btnh) # [B, T, K, H] + + value_proj = self.v_proj(x).reshape(b, t, self.num_kv_heads, self.head_dim) + value_proj = shard(value_proj, self.shd_cfg.act_btnh) # [B, T, K, H] + + # RoPE and Cache Logic + left_pads = count_left_pads(segment_ids) + left_pads = shard(left_pads, P(self.shd_cfg.act_btnh[0])) + cache.start_ind.value = jnp.where(cache.start_ind.value < 0, left_pads, cache.start_ind.value) + position_ids = compute_positions_from_segment_ids(segment_ids) + cache.cur_ind.value + sin, cos = _generate_pos_embeddings(position_ids, self.head_dim, self.rope_theta) + query_proj = apply_rope(query_proj, sin, cos) + key_proj = apply_rope(key_proj, sin, cos) + + # Ensure dtype matches cache and preserve sharding + # astype can break sharding, so re-shard after dtype conversion + cache_dtype = cache.k_cache.dtype + value_proj = shard(value_proj.astype(cache_dtype), self.shd_cfg.act_btnh) + key_proj = shard(key_proj.astype(cache_dtype), self.shd_cfg.act_btnh) + + # Update K/V cache [B, S, K, H] + slice_indices = (0, cache.cur_ind.value, 0, 0) + cache.v_cache.value = jax.lax.dynamic_update_slice(cache.v_cache.value, value_proj, slice_indices) + cache.k_cache.value = jax.lax.dynamic_update_slice(cache.k_cache.value, key_proj, slice_indices) + + b, t, n, h = query_proj.shape + + # GQA reshape and attention logits + query_proj_gqa = query_proj.reshape((b, t, self.num_kv_heads, self.n_rep, h)) + attn_logits = jnp.einsum("BTKGH,BSKH->BTSKG", query_proj_gqa, cache.k_cache.value) * self.scale + + # Masking and Softmax + q_pos = cache.cur_ind.value + jnp.arange(t, dtype=jnp.int32)[None, :] - cache.start_ind.value[:, None] + ts = jnp.arange(cache.size, dtype=jnp.int32) # (cache.size,) + kv_segment_ids = (ts[None, :] >= cache.start_ind.value[:, None]) & (ts[None, :] < cache.cur_ind.value + t) + k_pos = ts[None, :] - cache.start_ind.value[:, None] # (b, cache.size) + + # Segment mask (always applied) + segment_mask = kv_segment_ids[:, None, :] == segment_ids[:, :, None] + + # Conditionally apply causal masking + if self.cfg.use_causal_mask: + causal_mask = k_pos[:, None, :] <= q_pos[:, :, None] + final_mask = causal_mask & segment_mask # (B, T, S) + else: + # Bidirectional attention: only use segment mask + final_mask = segment_mask # (B, T, S) + + attn_mask = final_mask[:, :, :, None, None] + attn_logits = jnp.where(attn_mask, attn_logits, _K_MASK) + + # Softmax + attn_weights = jax.nn.softmax(attn_logits.astype(jnp.float32), axis=2).astype(attn_logits.dtype) + qkv = jnp.einsum("BTSKG,BSKH->BTKGH", attn_weights, cache.v_cache.value) + qkv = qkv.reshape((b, t, n, h)) + + # Reshape for o_proj: [B, T, N, H] -> [B, T, N*H] + qkv_flat = qkv.reshape(b, t, n * h) + output = self.o_proj(qkv_flat) + + cache.cur_ind.value = cache.cur_ind.value + t + return shard(output, self.shd_cfg.act_btd) + + @property + def head_dim(self): + return self.cfg.head_dim + + @property + def num_heads(self): + return self.cfg.num_heads + + @property + def num_kv_heads(self): + return self.cfg.num_kv_heads + + +class DecoderLayer(nnx.Module): + def __init__(self, cfg: ModelConfig, *, rngs: nnx.Rngs): + self.input_layernorm = RMSNorm(cfg.emb_dim, cfg, rngs=rngs) + self.attn = Attention(cfg=cfg, rngs=rngs) + self.post_attention_layernorm = RMSNorm(cfg.emb_dim, cfg, rngs=rngs) + self.mlp = MLP(cfg=cfg, rngs=rngs) + + def __call__(self, x: Array, cache: LayerCache | None, segment_ids: Array) -> Array: + inputs_normalized = self.input_layernorm(x) + attn_output = x + self.attn(inputs_normalized, cache, segment_ids) + outputs = attn_output + self.mlp(self.post_attention_layernorm(attn_output)) + return outputs + + +class Qwen2(nnx.Module): + def __init__(self, cfg: ModelConfig, *, rngs: nnx.Rngs): + self.shd_cfg = cfg.shd_cfg + self.embedder = shard( + nnx.Embed(num_embeddings=cfg.vocab_size, features=cfg.emb_dim, dtype=jnp.bfloat16, rngs=rngs), + cfg.shd_cfg.emb_vd, + ) + self.out_emb_shd = None if get_abstract_mesh().empty else cfg.shd_cfg.act_btd + self.layers = nnx.List([DecoderLayer(cfg=cfg, rngs=rngs) for _ in range(cfg.num_layers)]) + self.final_norm = RMSNorm(cfg.emb_dim, cfg, rngs=rngs) + # Standard Linear layer for lm_head + self.lm_head = shard( + nnx.Linear(cfg.emb_dim, cfg.vocab_size, use_bias=False, rngs=rngs), + cfg.shd_cfg.emb_dv + ) + + def init_cache( + self, cfg: ModelConfig, batch_size: int, token_len: int, generate_steps: int, dtype: jnp.dtype = jnp.bfloat16 + ) -> Cache: + cache_size = 2 ** math.ceil(math.log2(max(token_len + generate_steps, 1))) + return [LayerCache(cfg, batch_size, cache_size, dtype) for _ in range(cfg.num_layers)] + + def __call__(self, tokens, segment_ids, cache, num_right_pads): + x = self.embedder.embedding.value.at[(tokens,)].get(out_sharding=self.out_emb_shd) + for i, layer in enumerate(self.layers): + x = layer(x, cache[i], segment_ids) + logits = self.lm_head(self.final_norm(x)) + + # For generation/sampling, replicate all dimensions across devices + # This will trigger automatic all-gather to prepare for sampling + if not get_abstract_mesh().empty: + # logits shape: [B, T, V], replicate all dims for sampling compatibility + logits = shard(logits, P(None, None, None)) + + return logits + + +@jax.jit +def forward(model: nnx.Module, cache: Cache, tokens: Array, pad_id: int) -> tuple[Array, nnx.Cache]: + segment_ids = 1 * (tokens != pad_id) + num_right_pads = count_right_pads(tokens, pad_id) + logits = model(tokens, segment_ids, cache, num_right_pads) + target_ind = tokens.shape[-1] - num_right_pads - 1 + return logits[:, target_ind], cache \ No newline at end of file diff --git a/bonsai/models/qwen2/params.py b/bonsai/models/qwen2/params.py new file mode 100644 index 00000000..5230f334 --- /dev/null +++ b/bonsai/models/qwen2/params.py @@ -0,0 +1,146 @@ +import gc +import re +from dataclasses import dataclass +from typing import Any + +import jax +import safetensors +from etils import epath +from flax import nnx + +from bonsai.models.qwen2 import modeling as model_lib + + +@dataclass(frozen=True) +class Transform: + permute: tuple[int, ...] | None = None + reshape: tuple[int, ...] | None = None + reshape_first: bool = False + + +TRANSFORM_LINEAR = Transform(permute=(1, 0)) +TRANSFORM_NONE = Transform() + + +def _get_key_and_transform_mapping(cfg: model_lib.ModelConfig) -> dict[str, tuple[str | None, Transform | None]]: + # For Linear layers, we only need simple transpose from PyTorch's (out, in) to JAX's (in, out) + + return { + r"model\.embed_tokens\.weight": ("embedder.embedding", TRANSFORM_NONE), + # Attention projections: simple transpose for Linear layers + r"model\.layers\.([0-9]+)\.self_attn\.q_proj\.weight": (r"layers.\1.attn.q_proj.kernel", TRANSFORM_LINEAR), + r"model\.layers\.([0-9]+)\.self_attn\.k_proj\.weight": (r"layers.\1.attn.k_proj.kernel", TRANSFORM_LINEAR), + r"model\.layers\.([0-9]+)\.self_attn\.v_proj\.weight": (r"layers.\1.attn.v_proj.kernel", TRANSFORM_LINEAR), + r"model\.layers\.([0-9]+)\.self_attn\.o_proj\.weight": (r"layers.\1.attn.o_proj.kernel", TRANSFORM_LINEAR), + # Attention biases: no transformation needed + r"model\.layers\.([0-9]+)\.self_attn\.q_proj\.bias": (r"layers.\1.attn.q_proj.bias", TRANSFORM_NONE), + r"model\.layers\.([0-9]+)\.self_attn\.k_proj\.bias": (r"layers.\1.attn.k_proj.bias", TRANSFORM_NONE), + r"model\.layers\.([0-9]+)\.self_attn\.v_proj\.bias": (r"layers.\1.attn.v_proj.bias", TRANSFORM_NONE), + # MLP projections + r"model\.layers\.([0-9]+)\.mlp\.gate_proj\.weight": (r"layers.\1.mlp.gate_proj.kernel", TRANSFORM_LINEAR), + r"model\.layers\.([0-9]+)\.mlp\.up_proj\.weight": (r"layers.\1.mlp.up_proj.kernel", TRANSFORM_LINEAR), + r"model\.layers\.([0-9]+)\.mlp\.down_proj\.weight": (r"layers.\1.mlp.down_proj.kernel", TRANSFORM_LINEAR), + # Normalization layers + r"model\.norm\.weight": ("final_norm.scale", TRANSFORM_NONE), + r"model\.layers\.([0-9]+)\.input_layernorm\.weight": (r"layers.\1.input_layernorm.scale", TRANSFORM_NONE), + r"model\.layers\.([0-9]+)\.post_attention_layernorm\.weight": ( + r"layers.\1.post_attention_layernorm.scale", + TRANSFORM_NONE, + ), + # LM head + r"lm_head\.weight": ("lm_head.kernel", TRANSFORM_LINEAR), + } + + +def _get_jax_key( + mapping: dict[str, tuple[str | None, Transform | None]], source_key: str +) -> tuple[str | None, Transform | None]: + for pat, (repl, transform) in mapping.items(): + if re.match(pat, source_key): + if repl is None: + return None, None + return re.sub(pat, repl, source_key), transform + + print(f"Warning: No mapping found for key '{source_key}', skipping...") + return None, None + + +def _assign_weights( + keys: list[str | int], + tensor: Any, + state_dict: dict, + st_key: str, + transform: Transform | None, + sharding_dict: dict | None, +) -> None: + key, *rest = keys + if not rest: + if transform is not None: + if transform.reshape_first and transform.reshape is not None: + tensor = tensor.reshape(transform.reshape) + if transform.permute is not None: + tensor = tensor.transpose(transform.permute) + if not transform.reshape_first and transform.reshape is not None: + tensor = tensor.reshape(transform.reshape) + + if tensor.shape != state_dict[key].shape: + raise ValueError(f"Shape mismatch for {st_key}: {tensor.shape} vs {state_dict[key].shape}") + + state_dict[key] = jax.device_put(tensor, sharding_dict[key] if sharding_dict else None) + else: + next_sharding = sharding_dict[key] if sharding_dict is not None else None + _assign_weights(rest, tensor, state_dict[key], st_key, transform, next_sharding) + + +def _stoi(s: str) -> str | int: + try: + return int(s) + except ValueError: + return s + + +def create_model_from_safe_tensors( + file_dir: str, cfg: model_lib.ModelConfig, mesh: jax.sharding.Mesh | None = None +) -> model_lib.Qwen2: + files = list(epath.Path(file_dir).expanduser().glob("*.safetensors")) + if not files: + raise ValueError(f"No safetensors found in {file_dir}") + + qwen2 = nnx.eval_shape(lambda: model_lib.Qwen2(cfg, rngs=nnx.Rngs(params=0))) + graph_def, abs_state = nnx.split(qwen2) + state_dict = abs_state.to_pure_dict() + sharding = nnx.get_named_sharding(abs_state, mesh).to_pure_dict() if mesh is not None else None + + key_mapping = _get_key_and_transform_mapping(cfg) + conversion_errors = [] + + for f in files: + with safetensors.safe_open(f, framework="numpy") as sf: + for torch_key in sf.keys(): + jax_key, transform = _get_jax_key(key_mapping, torch_key) + if jax_key is None: + continue + + keys = [_stoi(k) for k in jax_key.split(".")] + try: + tensor = sf.get_tensor(torch_key) + _assign_weights(keys, tensor, state_dict, torch_key, transform, sharding) + except Exception as e: + full_jax_key = ".".join([str(k) for k in keys]) + conversion_errors.append( + f"Failed to assign '{torch_key}' to '{full_jax_key}': {type(e).__name__}: {e}" + ) + gc.collect() + + if conversion_errors: + raise RuntimeError( + f"Encountered {len(conversion_errors)} weight conversion errors:\n" + "\n".join(conversion_errors) + ) + + if cfg.tie_word_embeddings: + state_dict["lm_head"]["kernel"] = state_dict["embedder"]["embedding"].T + + model = nnx.merge(graph_def, state_dict) + gc.collect() + + return model \ No newline at end of file From 25bec119fd7196ea3fb51497f73e33294da368ba Mon Sep 17 00:00:00 2001 From: haibo Date: Mon, 19 Jan 2026 12:57:31 +0800 Subject: [PATCH 02/18] doc(qwen2): add readme --- bonsai/models/qwen2/README.md | 34 ++++++++++++++++++++++++++++++++++ 1 file changed, 34 insertions(+) create mode 100644 bonsai/models/qwen2/README.md diff --git a/bonsai/models/qwen2/README.md b/bonsai/models/qwen2/README.md new file mode 100644 index 00000000..d2cf47f2 --- /dev/null +++ b/bonsai/models/qwen2/README.md @@ -0,0 +1,34 @@ +# Qwen2 in JAX + +This directory contains a pure JAX implementation of the [Qwen2 language model](https://qwenlm.github.io/blog/qwen2/), using the [Flax NNX](https://flax.readthedocs.io/en/v0.8.3/experimental/nnx/index.html) API. + +## Model Configuration Support Status + +| Model Name | Config Support Status | +| :--- | :--- | +| **Dense Models** | | +| [Qwen2-0.5B](https://huggingface.co/Qwen/Qwen2-0.5B) | **✅ Supported** | +| [Qwen2-1.5B](https://huggingface.co/Qwen/Qwen2-1.5B) | **✅ Supported** | +| [Qwen2-7B](https://huggingface.co/Qwen/Qwen2-7B) | **✅ Supported** | +| [Qwen2-72B](https://huggingface.co/Qwen/Qwen2-72B) | **✅ Supported** | + + +### Running this model + +Run Qwen2 in action, implemented in pure JAX. + +```sh +python3 -m bonsai.models.qwen2.tests.run_model +``` + + +## Usage Example + +## Model Configurations + +The implementation supports all Qwen2 model sizes: + +- **0.5B**: 24 layers, 896 hidden size, 14 attention heads, 2 key-value heads +- **1.5B**: 28 layers, 1536 hidden size, 12 attention heads, 2 key-value heads +- **7B**: 28 layers, 3584 hidden size, 28 attention heads, 4 key-value heads +- **72B**: 80 layers, 8192 hidden size, 64 attention heads, 8 key-value heads \ No newline at end of file From 570a46990594a3a297944380ebb97659dbb266e1 Mon Sep 17 00:00:00 2001 From: haibo Date: Mon, 19 Jan 2026 12:59:53 +0800 Subject: [PATCH 03/18] test: qwen2 --- bonsai/models/qwen2/tests/__init__.py | 13 +++ bonsai/models/qwen2/tests/run_model.py | 111 +++++++++++++++++++++++++ 2 files changed, 124 insertions(+) create mode 100644 bonsai/models/qwen2/tests/__init__.py create mode 100644 bonsai/models/qwen2/tests/run_model.py diff --git a/bonsai/models/qwen2/tests/__init__.py b/bonsai/models/qwen2/tests/__init__.py new file mode 100644 index 00000000..2aae938f --- /dev/null +++ b/bonsai/models/qwen2/tests/__init__.py @@ -0,0 +1,13 @@ +# Copyright 2025 The JAX Authors. +# +# 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. \ No newline at end of file diff --git a/bonsai/models/qwen2/tests/run_model.py b/bonsai/models/qwen2/tests/run_model.py new file mode 100644 index 00000000..93519437 --- /dev/null +++ b/bonsai/models/qwen2/tests/run_model.py @@ -0,0 +1,111 @@ +import jax +import os +import jax.numpy as jnp +import numpy as np +from jax import P +from jax._src.mesh import AxisType +from transformers import AutoTokenizer + +from bonsai.models.qwen2 import modeling, params +from bonsai.utils import GreedySampler, Sampler + + +def tokenize(tokenizer, input: list[str], shd: P | None = None): + pad_idx = tokenizer.pad_token_id + lines = [ + tokenizer.apply_chat_template( + [{"role": "user", "content": l}], tokenize=False, add_generation_prompt=True + ) + for l in input + ] + lines = [tokenizer.encode(line) for line in lines] + max_len = max(len(line) for line in lines) # Right-align, left-padding to the max token length. + return jnp.array([np.pad(l, (max_len - len(l), 0), constant_values=pad_idx) for l in lines], out_sharding=shd) + + +def run_model(): + + model_ckpt_path = os.path.expanduser("~/.cache/modelscope/hub/models/Qwen/Qwen2-7B") + + # Disable sharding - run on single GPU + config = modeling.ModelConfig.qwen2_7b(use_sharding=False) + # mesh, batch_shd = None, None + + mesh = jax.make_mesh((1, 4), ("fsdp", "tp"), axis_types=(AxisType.Explicit, AxisType.Explicit)) + batch_shd = P("fsdp", None) + jax.set_mesh(mesh) + + query = [ + "why sky is blue?", + ] + + tokenizer = AutoTokenizer.from_pretrained( + model_ckpt_path, + local_files_only=True, + trust_remote_code=True, + ) + + print(f"Tokenizer special tokens:") + print(f" eos_token: {tokenizer.eos_token} (ID: {tokenizer.eos_token_id})") + print(f" pad_token: {tokenizer.pad_token} (ID: {tokenizer.pad_token_id})") + + print() + tokens = tokenize(tokenizer, query, batch_shd) + batch_size, token_len = tokens.shape + + generate_steps = 1024 + print(f"\nLoading model...") + model = params.create_model_from_safe_tensors(model_ckpt_path, config, mesh) + print(f"Model loaded successfully\n") + + cache = model.init_cache(config, batch_size, token_len, generate_steps) + + key = jax.random.key(0) + sampler = Sampler(temperature=0.7, top_p=0.9, top_k=20) + jit_sampler = jax.jit(sampler) + + logits, cache = modeling.forward(model, cache, tokens, tokenizer.pad_token_id) + + tokens_list = [] + finished = jnp.zeros((batch_size,), dtype=jnp.bool_) + + im_end_token_id = tokenizer.encode("<|im_end|>")[0] + for i in range(generate_steps): + # CRITICAL: Split key for each step to avoid deterministic sampling + key, subkey = jax.random.split(key) + next_tokens = jit_sampler(logits, key=subkey) + + current_token_id = int(next_tokens.squeeze(-1)[0]) + + # Only check for actual EOS token + is_eos = (next_tokens.squeeze(-1) == tokenizer.eos_token_id) + is_im_end = (next_tokens.squeeze(-1) == im_end_token_id) + + + finished = finished | is_eos | is_im_end + + tokens_list.append(next_tokens) + + if finished.all(): + print(f"✓ Generation stopped at step {i+1}/{generate_steps} (EOS token reached)") + break + + # Continue generation + logits, cache = modeling.forward(model, cache, next_tokens, tokenizer.pad_token_id) + + all_output_tokens = jax.device_get(jnp.concatenate(tokens_list, axis=-1)) + for i, q in enumerate(query): + print(f"User:\n {q}") + seq_tokens = all_output_tokens[i] + eos_idx = np.where(seq_tokens == tokenizer.eos_token_id)[0] + if eos_idx.size > 0: + seq_tokens = seq_tokens[: eos_idx[0]] + decoded = tokenizer.decode(seq_tokens, skip_special_tokens=True) + print(f"Answer ({len(seq_tokens)} tokens):\n {decoded}\n\n") + + +if __name__ == "__main__": + run_model() + + +__all__ = ["run_model"] \ No newline at end of file From 8345bd1634ac23e7dff1812593ed95b8ac11528e Mon Sep 17 00:00:00 2001 From: haibo Date: Mon, 19 Jan 2026 13:00:42 +0800 Subject: [PATCH 04/18] mimo-audio support --- bonsai/models/mimo_audio/__init__.py | 21 + .../mimo_audio/mimo_audio_configuration.py | 118 +++++ bonsai/models/mimo_audio/modeling.py | 425 ++++++++++++++++++ bonsai/models/mimo_audio/params.py | 145 ++++++ 4 files changed, 709 insertions(+) create mode 100644 bonsai/models/mimo_audio/__init__.py create mode 100644 bonsai/models/mimo_audio/mimo_audio_configuration.py create mode 100644 bonsai/models/mimo_audio/modeling.py create mode 100644 bonsai/models/mimo_audio/params.py diff --git a/bonsai/models/mimo_audio/__init__.py b/bonsai/models/mimo_audio/__init__.py new file mode 100644 index 00000000..aefeeb45 --- /dev/null +++ b/bonsai/models/mimo_audio/__init__.py @@ -0,0 +1,21 @@ +from bonsai.models.mimo_audio.mimo_audio import MimoAudio +from bonsai.models.mimo_audio.modeling import ( + MiMoAudioConfig, + MiMoAudioArguments, + FlaxMiMoAudioForCausalLM, +) +from bonsai.models.mimo_audio.mimo_audio_tokenizer import ( + FlaxMiMoAudioTokenizer, + MiMoAudioTokenizerConfig, + MelSpectrogram, +) + +__all__ = [ + "MimoAudio", + "MiMoAudioConfig", + "MiMoAudioArguments", + "FlaxMiMoAudioForCausalLM", + "FlaxMiMoAudioTokenizer", + "MiMoAudioTokenizerConfig", + "MelSpectrogram", +] \ No newline at end of file diff --git a/bonsai/models/mimo_audio/mimo_audio_configuration.py b/bonsai/models/mimo_audio/mimo_audio_configuration.py new file mode 100644 index 00000000..87fe4389 --- /dev/null +++ b/bonsai/models/mimo_audio/mimo_audio_configuration.py @@ -0,0 +1,118 @@ +from dataclasses import dataclass +from typing import Optional, TYPE_CHECKING +from bonsai.models.qwen3.modeling import ShardingCfg + +if TYPE_CHECKING: + from bonsai.models.qwen2.modeling import ModelConfig as Qwen2Config + + +@dataclass +class MiMoAudioConfig: + + # 主 Transformer 配置 + vocab_size: int = 151680 + hidden_size: int = 4096 + num_hidden_layers: int = 36 + num_attention_heads: int = 32 + num_key_value_heads: int = 8 + intermediate_size: int = 11008 + max_position_embeddings: int = 8192 + rope_theta: int = 640000 + head_dim: int = 128 + + group_size: int = 4 + audio_channels: int = 8 + + # Local Transformer config + local_dim: int = 1024 + local_layers: int = 16 + local_attn_heads: int = 64 + local_ffn_dim: int = 4096 + local_attn_dropout: float = 0.1 + + # Input Local Transformer config + input_local_layers: int = 6 + input_local_dim: int = 1024 + input_full_attention: bool = True + + # Sharding config + shd_cfg: ShardingCfg = ShardingCfg.no_sharding() + + @classmethod + def with_sharding(cls, **kwargs): + kwargs['shd_cfg'] = ShardingCfg.default() + return cls(**kwargs) + + def create_qwen2_config(self) -> "Qwen2Config": + from bonsai.models.qwen2.modeling import ModelConfig as Qwen2Config + + return Qwen2Config( + num_layers=36, + vocab_size=151680, + emb_dim=4096, + mlp_dim=11008, + num_heads=32, + head_dim=128, + num_kv_heads=8, + rope_theta=640000, + norm_eps=1e-6, + tie_word_embeddings=False, + shd_cfg=self.shd_cfg, + ) + + def create_local_qwen2_config(self) -> "Qwen2Config": + from bonsai.models.qwen2.modeling import ModelConfig as Qwen2Config + + return Qwen2Config( + num_layers=16, + vocab_size=151680, + emb_dim=1024, + mlp_dim=4096, + num_heads=64, + head_dim=16, # 1024 // 64 = 16 + num_kv_heads=64, + rope_theta=640000, + norm_eps=1e-6, + tie_word_embeddings=False, + shd_cfg=self.shd_cfg, + ) + + def create_input_local_qwen2_config(self) -> "Qwen2Config": + from bonsai.models.qwen2.modeling import ModelConfig as Qwen2Config + + # input_full_attention=True -> use_causal_mask=False (bidirectional attention) + return Qwen2Config( + num_layers=6, + vocab_size=151680, + emb_dim=1024, + mlp_dim=4096, # 1024 * 4 + num_heads=64, + head_dim=16, # 1024 // 64 = 16 + num_kv_heads=64, + rope_theta=640000, + norm_eps=1e-6, + tie_word_embeddings=False, + use_causal_mask=False, + shd_cfg=self.shd_cfg, + ) + + +@dataclass +class MiMoAudioArguments: + """Arguments for special token indices""" + model_name_or_path: str + sosp_idx: int + eosp_idx: int + sostm_idx: int + eostm_idx: int + eot_idx: int + empty_idx: int + + +@dataclass +class MiMoSamplerConfig: + """Sampler configuration for text/audio generation""" + do_sample: bool = True + temperature: float = 1.0 + top_k: int = 50 + top_p: float = 0.95 diff --git a/bonsai/models/mimo_audio/modeling.py b/bonsai/models/mimo_audio/modeling.py new file mode 100644 index 00000000..3d6196cc --- /dev/null +++ b/bonsai/models/mimo_audio/modeling.py @@ -0,0 +1,425 @@ +from typing import Optional, Tuple, List +import jax +import jax.numpy as jnp +from flax import nnx +from bonsai.models.qwen2.modeling import Qwen2, ModelConfig as Qwen2Config, Cache +from bonsai.models.qwen3.modeling import shard +from bonsai.utils.samplers import Sampler, GreedySampler +from bonsai.models.mimo_audio.mimo_audio_configuration import ( + MiMoAudioConfig, + MiMoAudioArguments, + MiMoSamplerConfig +) + + +class MiMoSampler: + """Sampling utilities for generation""" + + def __init__(self, config: MiMoSamplerConfig): + self.config = config + if config.do_sample: + self._sampler = Sampler( + temperature=config.temperature, + top_k=config.top_k, + top_p=config.top_p + ) + else: + self._sampler = GreedySampler() + + def sample( + self, + logits: jnp.ndarray, + key: jax.random.PRNGKey, + removed_tokens: Optional[List[int]] = None + ) -> jnp.ndarray: + if removed_tokens: + for t in removed_tokens: + logits = logits.at[:, t].set(-jnp.inf) + + result = self._sampler(logits, key=key) # [B, 1] + return result[:, 0] # [B] + + +class FlaxMiMoAudioForCausalLM(nnx.Module): + def __init__( + self, + config: MiMoAudioConfig, + args: MiMoAudioArguments, + rngs: Optional[nnx.Rngs] = None, + dtype: jnp.dtype = jnp.bfloat16, + ): + if rngs is None: + rngs = nnx.Rngs(0) + + self.config = config + self.args = args + self.dtype = dtype + self.shd_cfg = config.shd_cfg + + # Fixed model-specific configurations + self.speech_vocab_sizes = [1025, 1025, 129, 129, 129, 129, 129, 129] + self.speech_empty_ids = [1024, 1024, 128, 128, 128, 128, 128, 128] + self.delay_pattern = [0, 1, 2, 3, 4, 5, 6, 7] + + self.group_size = config.group_size + self.audio_channels = config.audio_channels + + self.qwen2_config = config.create_qwen2_config() + self.local_qwen2_config = config.create_local_qwen2_config() + self.input_local_qwen2_config = config.create_input_local_qwen2_config() + + self.model = Qwen2(self.qwen2_config, rngs=rngs) + self.local_transformer = Qwen2(self.local_qwen2_config, rngs=rngs) + self.input_local_transformer = Qwen2(self.input_local_qwen2_config, rngs=rngs) + + self.local_transformer.embedder = None + self.input_local_transformer.embedder = None + + self.lm_head = shard( + nnx.Linear( + config.hidden_size, + config.vocab_size, + use_bias=False, + dtype=self.dtype, + rngs=rngs + ), + self.shd_cfg.emb_dv + ) + + self.local_transformer_lm_heads = nnx.List([ + shard( + nnx.Linear( + config.local_dim, + self.speech_vocab_sizes[i], + use_bias=False, + dtype=self.dtype, + rngs=rngs + ), + self.shd_cfg.emb_dv + ) + for i in range(self.audio_channels) + ]) + + self.speech_embeddings = nnx.List([ + shard( + nnx.Embed( + self.speech_vocab_sizes[i], + config.input_local_dim, + dtype=self.dtype, + rngs=rngs + ), + self.shd_cfg.emb_vd + ) + for i in range(self.audio_channels) + ]) + + self.speech_group_downcast = shard( + nnx.Linear( + config.input_local_dim * config.group_size, + config.hidden_size, + use_bias=False, + dtype=self.dtype, + rngs=rngs + ), + self.shd_cfg.ffw_weight_df + ) + + self.hidden_states_downcast = shard( + nnx.Linear( + config.hidden_size, + config.local_dim, + use_bias=False, + dtype=self.dtype, + rngs=rngs + ), + self.shd_cfg.ffw_weight_df + ) + + def apply_input_local_transformer( + self, + speech_embeddings: jnp.ndarray, + cache: Optional[Cache] = None + ) -> jnp.ndarray: + """Apply input local transformer to speech embeddings""" + B, T_groups, group_size, hidden_size = speech_embeddings.shape + + input_embeddings = speech_embeddings.reshape(B * T_groups, group_size, hidden_size) + segment_ids = jnp.ones((B * T_groups, group_size), dtype=jnp.int32) + + if cache is None: + cache = self.input_local_transformer.init_cache( + self.input_local_qwen2_config, + B * T_groups, + group_size, + generate_steps=0, + dtype=self.dtype + ) + + x = input_embeddings + for i, layer in enumerate(self.input_local_transformer.layers): + x = layer(x, cache[i], segment_ids) + x = self.input_local_transformer.final_norm(x) + + return x.reshape(B, T_groups, group_size, hidden_size) + + def _prepare_input_embeds( + self, + input_ids: jnp.ndarray, + text_embed_fn + ) -> jnp.ndarray: + """Prepare input embeddings from interleaved text and speech tokens""" + B = input_ids.shape[0] + + text_input_ids = input_ids[:, 0, ::self.group_size] + speech_input_ids = input_ids[:, 1:, :].reshape( + B, self.audio_channels, -1, self.group_size + ).transpose(0, 2, 1, 3) + + is_speech = text_input_ids == self.args.empty_idx + + speech_embeds = jnp.zeros( + (B, is_speech.shape[1], self.group_size, self.config.input_local_dim), + dtype=self.dtype + ) + + for idx in range(self.audio_channels): + cur_empty = self.speech_empty_ids[idx] + cur_embed = self.speech_embeddings[idx] + cur_speech_ids = speech_input_ids[:, :, idx, :] + cur_speech_embeds = cur_embed(cur_speech_ids) + + cur_mask = cur_speech_ids == cur_empty + cur_speech_embeds = cur_speech_embeds * ~cur_mask[..., None] + speech_embeds = speech_embeds + cur_speech_embeds + + speech_embeds = speech_embeds * is_speech[:, :, None, None] + speech_embeds = self.apply_input_local_transformer(speech_embeds, cache=None) + + speech_embeds = speech_embeds * is_speech[:, :, None, None] + + T_groups = speech_embeds.shape[1] + speech_grouped_embeds = self.speech_group_downcast( + speech_embeds.reshape(B, T_groups, -1) + ) + + text_input_ids_safe = jnp.where(text_input_ids == -100, 0, text_input_ids) + text_embeds = text_embed_fn(text_input_ids_safe) + + text_zero_mask = (text_input_ids == self.args.empty_idx) | (text_input_ids == -100) + text_embeds = text_embeds * ~text_zero_mask[..., None] + + output = text_embeds + speech_grouped_embeds + return shard(output, self.shd_cfg.act_btd) + + def forward( + self, + input_ids: jnp.ndarray, + cache: Cache, + pad_id: int = 0, + ) -> Tuple[jnp.ndarray, jnp.ndarray, Cache]: + """Forward pass through the model""" + text_input_ids = input_ids[:, 0, ::self.group_size] + + def text_embed_fn(x): + return self.model.embedder.embedding.value[x] + + inputs_embeds = self._prepare_input_embeds(input_ids, text_embed_fn) + + B, T_groups, _ = inputs_embeds.shape + segment_ids = 1 * (text_input_ids != -100) + + # Run through main transformer + x = inputs_embeds + for i, layer in enumerate(self.model.layers): + x = layer(x, cache[i], segment_ids) + hidden_states = self.model.final_norm(x) # [B, T_groups, hidden_size] + + text_logits = self.lm_head(hidden_states[:, -1:, :]) # [B, 1, vocab_size] + + # Downcast hidden states for local transformer + local_hidden_states = self.hidden_states_downcast( + hidden_states[:, -1:, :] + ) # [B, 1, local_dim] + + return text_logits, local_hidden_states, cache + + def local_forward( + self, + local_embeds: jnp.ndarray, # [B, 1, local_dim] + key: jax.random.PRNGKey, + local_sampler: Optional[MiMoSampler] = None, + ) -> jnp.ndarray: + """ + Generate audio tokens for one group using local transformer. + + Args: + local_embeds: [B, 1, local_dim] + key: Random key for sampling + local_sampler: Sampler configuration + + Returns: + local_tokens: [B, group_size, audio_channels] + """ + B = local_embeds.shape[0] + delay_iters = self.group_size + max(self.delay_pattern) + + local_tokens = jnp.zeros( + (B, self.group_size, self.audio_channels), + dtype=jnp.int32 + ) + + if local_sampler is None: + local_sampler = MiMoSampler(MiMoSamplerConfig()) + + cache = self.local_transformer.init_cache( + self.local_qwen2_config, + B, + token_len=1, + generate_steps=delay_iters - 1, + dtype=self.dtype, + ) + + segment_ids = jnp.ones((B, 1), dtype=jnp.int32) + + for t in range(delay_iters): + hidden_state, cache = _local_transformer_step_jit( + self.local_transformer, local_embeds, cache, segment_ids + ) + + next_local_embeds = jnp.zeros_like(local_embeds) + + for idx in range(self.audio_channels): + cur_start = self.delay_pattern[idx] + cur_end = cur_start + self.group_size + cur_empty = self.speech_empty_ids[idx] + + if cur_start <= t < cur_end: + cur_lm_head = self.local_transformer_lm_heads[idx] + cur_logits = cur_lm_head(hidden_state[:, -1, :]) + + key, subkey = jax.random.split(key) + cur_token = local_sampler.sample( + cur_logits, + subkey, + removed_tokens=[cur_empty] + ) + + local_tokens = local_tokens.at[:, t - cur_start, idx].set(cur_token) + + cur_input_embed = self.speech_embeddings[idx](cur_token[:, None]) + + next_local_embeds = next_local_embeds + cur_input_embed + + local_embeds = next_local_embeds + + return local_tokens + + +# ============================================================================ +# JIT-compiled functions for fast inference +# ============================================================================ + +@jax.jit +def _local_transformer_step_jit( + local_transformer: nnx.Module, + local_embeds: jnp.ndarray, + cache: Cache, + segment_ids: jnp.ndarray, +) -> Tuple[jnp.ndarray, Cache]: + """ + JIT-compiled single step of local transformer forward pass. + + This is a helper function to accelerate the inner loop of local_forward. + Being a module-level function (not instance method) allows JAX to properly + JIT compile it. + + Args: + local_transformer: The local transformer module + local_embeds: [B, 1, local_dim] + cache: Cache for local transformer + segment_ids: [B, 1] + + Returns: + hidden_state: [B, 1, local_dim] + cache: Updated cache (IMPORTANT for correct behavior) + """ + x = local_embeds + for i, layer in enumerate(local_transformer.layers): + x = layer(x, cache[i], segment_ids) + hidden_state = local_transformer.final_norm(x) + return hidden_state, cache + + +@jax.jit +def forward_jit( + model: FlaxMiMoAudioForCausalLM, + input_ids: jnp.ndarray, + cache: Cache, + pad_id: int = 0, +) -> Tuple[jnp.ndarray, jnp.ndarray, Cache]: + """ + JIT-compiled forward pass for fast inference. + + Similar to qwen2's forward function, this returns the cache to enable + proper JAX tracing of stateful computations. + + Args: + model: FlaxMiMoAudioForCausalLM instance + input_ids: [B, audio_channels + 1, T * group_size] + cache: Cache for KV storage + pad_id: Padding token ID + + Returns: + text_logits: [B, 1, vocab_size] + local_hidden_states: [B, 1, local_dim] + cache: Updated cache (for JAX tracing) + """ + text_logits, local_hidden_states, cache = model.forward(input_ids, cache, pad_id) + return text_logits, local_hidden_states, cache + + +@jax.jit +def local_forward_jit( + model: FlaxMiMoAudioForCausalLM, + local_embeds: jnp.ndarray, + key: jax.random.PRNGKey, +) -> jnp.ndarray: + """ + JIT-compiled local forward pass for audio generation. + + NOTE: This version uses greedy sampling (no sampler parameter for JIT simplicity). + For temperature-based sampling, use model.local_forward() directly. + + Args: + model: FlaxMiMoAudioForCausalLM instance + local_embeds: [B, 1, local_dim] + key: Random key (used if needed in future, currently greedy) + + Returns: + audio_tokens: [B, group_size, audio_channels] + """ + # Use greedy sampling for JIT-compiled version + return model.local_forward(local_embeds, key, local_sampler=None) + + +# Example usage: +if __name__ == "__main__": + # Create configuration + config = MiMoAudioConfig() + args = MiMoAudioArguments( + model_name_or_path="mimo-audio", + sosp_idx=151646, + eosp_idx=151647, + sostm_idx=151648, + eostm_idx=151649, + eot_idx=151643, + empty_idx=151645, + ) + + # Create model + model = FlaxMiMoAudioForCausalLM(config,args) + + print("Model created successfully!") + print(f"Audio channels: {model.audio_channels}") + print(f"Group size: {model.group_size}") + print(f"Speech vocab sizes: {model.speech_vocab_sizes}") diff --git a/bonsai/models/mimo_audio/params.py b/bonsai/models/mimo_audio/params.py new file mode 100644 index 00000000..759707df --- /dev/null +++ b/bonsai/models/mimo_audio/params.py @@ -0,0 +1,145 @@ +import re +from typing import Any + +import jax +import jax.numpy as jnp +from flax import nnx + +from bonsai.models.qwen2.params import Transform, TRANSFORM_LINEAR, TRANSFORM_NONE, _stoi, _assign_weights + + +def _get_qwen2_key_mapping(prefix: str) -> dict[str, tuple[str, Transform]]: + q_bias_flatten = Transform(reshape=(-1,)) + kv_bias_flatten = Transform(reshape=(-1,)) + + return { + rf"{prefix}\.embed_tokens\.weight": (f"{prefix}.embedder.embedding", TRANSFORM_NONE), + rf"{prefix}\.layers\.([0-9]+)\.self_attn\.q_proj\.weight": (rf"{prefix}.layers.\1.attn.q_proj.kernel", TRANSFORM_LINEAR), + rf"{prefix}\.layers\.([0-9]+)\.self_attn\.k_proj\.weight": (rf"{prefix}.layers.\1.attn.k_proj.kernel", TRANSFORM_LINEAR), + rf"{prefix}\.layers\.([0-9]+)\.self_attn\.v_proj\.weight": (rf"{prefix}.layers.\1.attn.v_proj.kernel", TRANSFORM_LINEAR), + rf"{prefix}\.layers\.([0-9]+)\.self_attn\.o_proj\.weight": (rf"{prefix}.layers.\1.attn.o_proj.kernel", TRANSFORM_LINEAR), + rf"{prefix}\.layers\.([0-9]+)\.self_attn\.q_proj\.bias": (rf"{prefix}.layers.\1.attn.q_proj.bias", q_bias_flatten), + rf"{prefix}\.layers\.([0-9]+)\.self_attn\.k_proj\.bias": (rf"{prefix}.layers.\1.attn.k_proj.bias", kv_bias_flatten), + rf"{prefix}\.layers\.([0-9]+)\.self_attn\.v_proj\.bias": (rf"{prefix}.layers.\1.attn.v_proj.bias", kv_bias_flatten), + rf"{prefix}\.layers\.([0-9]+)\.mlp\.gate_proj\.weight": (rf"{prefix}.layers.\1.mlp.gate_proj.kernel", TRANSFORM_LINEAR), + rf"{prefix}\.layers\.([0-9]+)\.mlp\.up_proj\.weight": (rf"{prefix}.layers.\1.mlp.up_proj.kernel", TRANSFORM_LINEAR), + rf"{prefix}\.layers\.([0-9]+)\.mlp\.down_proj\.weight": (rf"{prefix}.layers.\1.mlp.down_proj.kernel", TRANSFORM_LINEAR), + rf"{prefix}\.norm\.weight": (f"{prefix}.final_norm.scale", TRANSFORM_NONE), + rf"{prefix}\.layers\.([0-9]+)\.input_layernorm\.weight": (rf"{prefix}.layers.\1.input_layernorm.scale", TRANSFORM_NONE), + rf"{prefix}\.layers\.([0-9]+)\.post_attention_layernorm\.weight": ( + rf"{prefix}.layers.\1.post_attention_layernorm.scale", + TRANSFORM_NONE, + ), + } + + +def _get_mimo_key_mapping(audio_channels: int) -> dict[str, tuple[str, Transform]]: + """Generate key mapping for MiMo-specific layers.""" + mapping = { + r"lm_head\.weight": ("lm_head.kernel", TRANSFORM_LINEAR), + r"hidden_states_downcast\.weight": ("hidden_states_downcast.kernel", TRANSFORM_LINEAR), + r"speech_group_downcast\.weight": ("speech_group_downcast.kernel", TRANSFORM_LINEAR), + } + + for i in range(audio_channels): + mapping[rf"speech_embeddings\.{i}\.weight"] = (f"speech_embeddings.{i}.embedding", TRANSFORM_NONE) + mapping[rf"local_transformer_lm_heads\.{i}\.weight"] = (f"local_transformer_lm_heads.{i}.kernel", TRANSFORM_LINEAR) + + return mapping + + +def _get_jax_key( + mapping: dict[str, tuple[str, Transform]], source_key: str +) -> tuple[str | None, Transform | None]: + """Get JAX key from source key using regex mapping.""" + for pat, (jax_key, transform) in mapping.items(): + match = re.fullmatch(pat, source_key) + if match: + result_key = re.sub(pat, jax_key, source_key) + return result_key, transform + return None, None + + +def create_model_with_weights( + model_path: str, + config, + args, + rngs: nnx.Rngs | None = None, + dtype: Any = jnp.bfloat16, + mesh: jax.sharding.Mesh | None = None, +) -> Any: + """Create MiMo Audio model and load weights from safetensors.""" + import os + import json + from safetensors import safe_open + + if rngs is None: + rngs = nnx.Rngs(0) + + print("Creating MiMo Audio model with weights...") + + from bonsai.models.mimo_audio.modeling import FlaxMiMoAudioForCausalLM + + model = nnx.eval_shape(lambda: FlaxMiMoAudioForCausalLM(config, args, nnx.Rngs(0), dtype=dtype)) + graph_def, abs_state = nnx.split(model) + pure_state_dict = abs_state.to_pure_dict() + + sharding = nnx.get_named_sharding(abs_state, mesh).to_pure_dict() if mesh is not None else None + + index_path = os.path.join(model_path, "model.safetensors.index.json") + state_dict = {} + + if os.path.exists(index_path): + with open(index_path) as f: + index = json.load(f) + for shard_file in sorted(set(index["weight_map"].values())): + with safe_open(os.path.join(model_path, shard_file), framework="numpy") as f: + for key in f.keys(): + state_dict[key] = f.get_tensor(key) + else: + safetensors_path = os.path.join(model_path, "model.safetensors") + with safe_open(safetensors_path, framework="numpy") as f: + for key in f.keys(): + state_dict[key] = f.get_tensor(key) + + full_mapping = {} + + # Main model mapping (model prefix) + full_mapping.update(_get_qwen2_key_mapping("model")) + + # Local transformer mapping + full_mapping.update(_get_qwen2_key_mapping("local_transformer")) + + # Input local transformer mapping + full_mapping.update(_get_qwen2_key_mapping("input_local_transformer")) + + full_mapping.update(_get_mimo_key_mapping(config.audio_channels)) + + conversion_errors = [] + loaded_count = 0 + + for torch_key, tensor in state_dict.items(): + jax_key, transform = _get_jax_key(full_mapping, torch_key) + if jax_key is None: + continue + + keys = [_stoi(k) for k in jax_key.split(".")] + try: + tensor_jax = jnp.asarray(tensor, dtype=dtype) + _assign_weights(keys, tensor_jax, pure_state_dict, torch_key, transform, sharding) + loaded_count += 1 + except KeyError as e: + continue + except Exception as e: + full_jax_key = ".".join([str(k) for k in keys]) + conversion_errors.append(f"Failed: '{torch_key}' -> '{full_jax_key}': {e}") + + if conversion_errors: + raise RuntimeError( + f"Encountered {len(conversion_errors)} weight conversion errors:\n" + "\n".join(conversion_errors[:10]) + ) + + model = nnx.merge(graph_def, pure_state_dict) + + print(f"MiMo Audio model created successfully! ({loaded_count} parameters loaded)") + return model From c8bc92e8d0dc629cf25b5899f9f4c37396d3ebb7 Mon Sep 17 00:00:00 2001 From: haibo Date: Mon, 19 Jan 2026 13:01:11 +0800 Subject: [PATCH 05/18] mimo-audio-tokenizer support --- .../models/mimo_audio/mimo_audio_tokenizer.py | 764 ++++++++++++++++++ .../mimo_audio_tokenizer_configuration.py | 227 ++++++ .../mimo_audio/mimo_audio_tokenizer_params.py | 249 ++++++ 3 files changed, 1240 insertions(+) create mode 100644 bonsai/models/mimo_audio/mimo_audio_tokenizer.py create mode 100644 bonsai/models/mimo_audio/mimo_audio_tokenizer_configuration.py create mode 100644 bonsai/models/mimo_audio/mimo_audio_tokenizer_params.py diff --git a/bonsai/models/mimo_audio/mimo_audio_tokenizer.py b/bonsai/models/mimo_audio/mimo_audio_tokenizer.py new file mode 100644 index 00000000..2b668b0d --- /dev/null +++ b/bonsai/models/mimo_audio/mimo_audio_tokenizer.py @@ -0,0 +1,764 @@ +import math +from typing import Optional, Sequence, Tuple +import jax +import jax.numpy as jnp +from flax import nnx +from jax.sharding import get_abstract_mesh, reshard + +from bonsai.models.mimo_audio.mimo_audio_tokenizer_configuration import ( + MiMoShardingCfg, + MiMoAudioTokenizerConfig, + EncoderOutput, + VocoderOutput, +) + +Array = jnp.ndarray + + +def shard(x: Array, s) -> Array: + """Apply sharding to array if mesh is available.""" + mesh = get_abstract_mesh() + if not mesh.empty and len(mesh.axis_names) > 0: + return reshard(x, s) + return x + + +def make_sequence_mask(lengths: Array, max_length: Optional[int] = None) -> Array: + max_len = max_length or int(jnp.max(lengths)) + base = jnp.arange(max_len)[None, :] + return base < lengths[:, None] + + +def get_position_ids(lengths: Array, max_length: Optional[int] = None) -> Array: + max_len = max_length or int(jnp.max(lengths)) + base = jnp.arange(max_len)[None, :] + return jnp.broadcast_to(base, (lengths.shape[0], max_len)) + + +def rotate_half(x: Array) -> Array: + x1, x2 = jnp.split(x, 2, axis=-1) + return jnp.concatenate([-x2, x1], axis=-1) + + +def apply_rotary(x: Array, cos: Array, sin: Array) -> Array: + cos = cos[:, None, :, :] + sin = sin[:, None, :, :] + return (x * cos) + (rotate_half(x) * sin) + + +class MelSpectrogram: + """Mel spectrogram computation for audio processing.""" + + def __init__( + self, + sample_rate: int, + n_fft: int, + hop_length: int, + win_length: int, + f_min: float, + f_max: float, + n_mels: int, + power: float = 1.0, + center: bool = True, + ) -> None: + self.sample_rate = int(sample_rate) + self.n_fft = int(n_fft) + self.hop_length = int(hop_length) + self.win_length = int(win_length) + self.f_min = float(f_min) + self.f_max = float(f_max) + self.n_mels = int(n_mels) + self.power = float(power) + self.center = center + + self._window = self._build_window() + self._mel_filterbank = self._build_mel_filterbank() + + def _build_window(self) -> jnp.ndarray: + if self.win_length <= 1: + return jnp.ones((self.win_length,), dtype=jnp.float32) + n = jnp.arange(self.win_length, dtype=jnp.float32) + return 0.5 - 0.5 * jnp.cos(2 * jnp.pi * n / (self.win_length - 1)) + + def _hz_to_mel(self, freq: jnp.ndarray) -> jnp.ndarray: + return 2595.0 * jnp.log10(1.0 + freq / 700.0) + + def _mel_to_hz(self, mel: jnp.ndarray) -> jnp.ndarray: + return 700.0 * (jnp.power(10.0, mel / 2595.0) - 1.0) + + def _build_mel_filterbank(self) -> jnp.ndarray: + freq_bins = jnp.linspace( + 0.0, + self.sample_rate / 2, + self.n_fft // 2 + 1, + dtype=jnp.float32, + ) + mel_min = self._hz_to_mel(self.f_min) + mel_max = self._hz_to_mel(self.f_max) + mel_points = jnp.linspace(mel_min, mel_max, self.n_mels + 2) + hz_points = self._mel_to_hz(mel_points) + + filterbanks = [] + for i in range(self.n_mels): + lower = hz_points[i] + center = hz_points[i + 1] + upper = hz_points[i + 2] + denom_left = jnp.maximum(center - lower, 1e-10) + denom_right = jnp.maximum(upper - center, 1e-10) + left_slope = (freq_bins - lower) / denom_left + right_slope = (upper - freq_bins) / denom_right + filterbanks.append(jnp.maximum(0.0, jnp.minimum(left_slope, right_slope))) + + return jnp.stack(filterbanks, axis=0) + + def _frame_signal(self, waveform: jnp.ndarray) -> jnp.ndarray: + frame_length = self.n_fft + + if self.center: + pad = self.n_fft // 2 + if waveform.shape[0] > 1: + waveform = jnp.pad(waveform, (pad, pad), mode="reflect") + else: + waveform = jnp.pad(waveform, (pad, pad)) + + total_length = int(waveform.shape[0]) + if total_length < frame_length: + pad_amount = frame_length - total_length + waveform = jnp.pad(waveform, (0, pad_amount)) + total_length = int(waveform.shape[0]) + + num_frames = 1 + max(0, (total_length - frame_length) // self.hop_length) + if num_frames <= 0: + num_frames = 1 + + starts = [idx * self.hop_length for idx in range(num_frames)] + frames = jnp.stack( + [waveform[start: start + frame_length] for start in starts], axis=0 + ) + return frames + + def _mel_spectrogram(self, waveform: jnp.ndarray) -> jnp.ndarray: + waveform = jnp.asarray(waveform, dtype=jnp.float32) + frames = self._frame_signal(waveform) + if self.win_length < self.n_fft: + total_pad = self.n_fft - self.win_length + pad_left = total_pad // 2 + pad_right = total_pad - pad_left + window = jnp.pad(self._window, (pad_left, pad_right)) + else: + window = self._window[: self.n_fft] + windowed = frames * window + stft = jnp.fft.rfft(windowed, n=self.n_fft, axis=1) + magnitude = jnp.abs(stft) ** self.power + mel_spec = magnitude @ self._mel_filterbank.T + return mel_spec.T + + def __call__(self, waveform: jnp.ndarray) -> jnp.ndarray: + """Compute mel spectrogram from waveform. + + Args: + waveform: JAX array of shape (samples,) or (batch, samples) + + Returns: + Mel spectrogram of shape (n_mels, time) or (batch, n_mels, time) + """ + waveform = jnp.asarray(waveform, dtype=jnp.float32) + + squeeze_dim = False + if waveform.ndim == 1: + waveform = waveform[jnp.newaxis, :] + squeeze_dim = True + + mel_outputs = [] + for sample in waveform: + mel_outputs.append(self._mel_spectrogram(sample)) + + mel_stack = jnp.stack(mel_outputs, axis=0) + if squeeze_dim: + mel_stack = jnp.squeeze(mel_stack, axis=0) + return mel_stack + + +class RotaryEmbedding(nnx.Module): + def __init__(self, base: float, dim: int, max_seq_len: int, rope_type: str = "default", dtype=jnp.float32): + self.base = base + self.dim = dim + self.max_seq_len = max_seq_len + self.rope_type = rope_type + self.dtype = dtype + half_dim = dim // 2 + inv_freq = 1.0 / (self.base ** (jnp.arange(0, half_dim, dtype=jnp.float32) / float(half_dim))) + self.inv_freq = nnx.Param(inv_freq) + self.attention_scaling = 1.0 + + def __call__(self, hidden_states: Array, position_ids: Array) -> Tuple[Array, Array]: + freq = position_ids[..., None] * self.inv_freq[None, None, :] + emb = jnp.concatenate([freq, freq], axis=-1) + cos = jnp.cos(emb) * self.attention_scaling + sin = jnp.sin(emb) * self.attention_scaling + return cos.astype(hidden_states.dtype), sin.astype(hidden_states.dtype) + + +class ConvTranspose1d(nnx.Module): + """Custom 1D transposed convolution for specific audio processing requirements.""" + + def __init__(self, in_channels: int, out_channels: int, kernel_size: int, stride: int, + shd_cfg: MiMoShardingCfg | None = None, + dtype=jnp.float32, rngs: Optional[nnx.Rngs] = None): + self.stride = stride + self.shd_cfg = shd_cfg or MiMoShardingCfg.no_sharding() + + kshape = (in_channels, out_channels, kernel_size) + kernel = jnp.zeros(kshape, dtype=dtype) + self.kernel = shard(nnx.Param(kernel), self.shd_cfg.conv_transpose_weight) + + bias = jnp.zeros((out_channels,), dtype=dtype) + self.bias = shard(nnx.Param(bias), self.shd_cfg.conv_transpose_bias) + + def __call__(self, x: Array) -> Array: + batch, length, channels = x.shape + kernel = self.kernel.value + kernel_size = kernel.shape[-1] + up_len = (length - 1) * self.stride + 1 + idx = jnp.arange(length) * self.stride + upsampled = jnp.zeros((batch, up_len, channels), dtype=x.dtype) + upsampled = upsampled.at[:, idx, :].set(x) + upsampled = jnp.pad(upsampled, ((0, 0), (kernel_size - 1, kernel_size - 1), (0, 0))) + lhs = jnp.swapaxes(upsampled, 1, 2) + rhs = jnp.flip(kernel, axis=-1).transpose(1, 0, 2) + y = jax.lax.conv_general_dilated( + lhs=lhs, + rhs=rhs, + window_strides=(1,), + padding='VALID', + dimension_numbers=('NCH', 'OIH', 'NCH'), + ) + y = y + self.bias.value[None, :, None] + y = jnp.swapaxes(y, 1, 2) + return y + + +class ISTFT(nnx.Module): + def __init__(self, n_fft: int, hop_length: int, win_length: int, padding: str = "same", + shd_cfg: MiMoShardingCfg | None = None, dtype=jnp.float32): + self.n_fft = n_fft + self.hop_length = hop_length + self.win_length = win_length + self.padding = padding + self.shd_cfg = shd_cfg or MiMoShardingCfg.no_sharding() + + self.window = shard( + nnx.Param(jnp.hanning(win_length).astype(dtype)), + self.shd_cfg.istft_window + ) + + self.pad = (self.win_length - self.hop_length) // 2 if padding == "same" else 0 + + def __call__(self, spec: Array) -> Array: + frames = jnp.fft.irfft(spec, n=self.n_fft, axis=1, norm="backward") + frames = frames * self.window[None, :, None] + frames = jnp.swapaxes(frames, 1, 2) + batch, num_frames, _ = frames.shape + output_size = (num_frames - 1) * self.hop_length + self.win_length + audio = jnp.zeros((batch, output_size), dtype=frames.dtype) + env = jnp.zeros_like(audio) + window_sq = jnp.square(self.window) + + def body(i, carry): + audio_acc, env_acc = carry + start = i * self.hop_length + frame = frames[:, i, :] + current_audio = jax.lax.dynamic_slice( + audio_acc, + (0, start), + (batch, self.win_length), + ) + current_env = jax.lax.dynamic_slice( + env_acc, + (0, start), + (batch, self.win_length), + ) + updated_audio = current_audio + frame + updated_env = current_env + window_sq + audio_acc = jax.lax.dynamic_update_slice(audio_acc, updated_audio, (0, start)) + env_acc = jax.lax.dynamic_update_slice(env_acc, updated_env, (0, start)) + return audio_acc, env_acc + + audio, env = jax.lax.fori_loop(0, num_frames, body, (audio, env)) + if self.pad > 0: + audio = audio[:, self.pad: -self.pad] + env = env[:, self.pad: -self.pad] + env = jnp.maximum(env, 1e-11) + audio = audio / env + return audio + + +class ISTFTHead(nnx.Module): + def __init__(self, dim: int, n_fft: int, hop_length: int, padding: str = "same", + shd_cfg: MiMoShardingCfg | None = None, + dtype=jnp.float32, rngs: Optional[nnx.Rngs] = None): + self.shd_cfg = shd_cfg or MiMoShardingCfg.no_sharding() + + self.linear = shard( + nnx.Linear(dim, n_fft + 2, dtype=dtype, rngs=rngs), + self.shd_cfg.istft_linear_weight + ) + + self.istft = ISTFT(n_fft=n_fft, hop_length=hop_length, win_length=n_fft, + padding=padding, shd_cfg=self.shd_cfg, dtype=dtype) + + def __call__(self, hidden_states: Array) -> Array: + x = self.linear(hidden_states) + x = jnp.swapaxes(x, 1, 2) + mag, phase = jnp.split(x, 2, axis=1) + + original_dtype = hidden_states.dtype + mag = mag.astype(jnp.float32) + phase = phase.astype(jnp.float32) + + mag = jnp.clip(jnp.exp(mag), a_max=1e2) + real = jnp.cos(phase) + imag = jnp.sin(phase) + spec = mag * (real + 1j * imag) + + audio = self.istft(spec) + audio = audio.astype(original_dtype) + return audio + + +class Attention(nnx.Module): + def __init__(self, embed_dim: int, num_heads: int, window_size: Tuple[int, int], causal: bool, + shd_cfg: MiMoShardingCfg, dtype=jnp.float32, + rngs: Optional[nnx.Rngs] = None): + self.embed_dim = embed_dim + self.num_heads = num_heads + self.head_dim = embed_dim // num_heads + self.scale = 1.0 / math.sqrt(self.head_dim) + self.window_size = window_size + self.causal = causal + self.shd_cfg = shd_cfg + + self.q_proj = shard( + nnx.Linear(embed_dim, embed_dim, use_bias=True, dtype=dtype, rngs=rngs), + shd_cfg.attn_qkvo_weight + ) + self.k_proj = shard( + nnx.Linear(embed_dim, embed_dim, use_bias=False, dtype=dtype, rngs=rngs), + shd_cfg.attn_qkvo_weight + ) + self.v_proj = shard( + nnx.Linear(embed_dim, embed_dim, use_bias=True, dtype=dtype, rngs=rngs), + shd_cfg.attn_qkvo_weight + ) + self.out_proj = shard( + nnx.Linear(embed_dim, embed_dim, dtype=dtype, rngs=rngs), + shd_cfg.attn_qkvo_weight + ) + + def _window_mask(self, seq_len: int) -> Optional[Array]: + left, right = self.window_size + if left < 0 and right < 0: + return None + pos = jnp.arange(seq_len) + rel = pos[None, :] - pos[:, None] + mask = jnp.ones((seq_len, seq_len), dtype=bool) + if left >= 0: + mask &= rel >= -left + if right >= 0: + mask &= rel <= right + return mask + + def __call__(self, x: Array, mask: Optional[Array], rope: Optional[Tuple[Array, Array]]) -> Array: + batch, seq_len, _ = x.shape + q = self.q_proj(x) + k = self.k_proj(x) + v = self.v_proj(x) + + def reshape(t): + t = t.reshape(batch, seq_len, self.num_heads, self.head_dim) + return jnp.swapaxes(t, 1, 2) + + q, k, v = reshape(q), reshape(k), reshape(v) + + q = shard(q, self.shd_cfg.act_btnh) + k = shard(k, self.shd_cfg.act_btnh) + v = shard(v, self.shd_cfg.act_btnh) + + if rope is not None: + cos, sin = rope + q = apply_rotary(q, cos, sin) + k = apply_rotary(k, cos, sin) + scores = jnp.einsum("bhqd,bhkd->bhqk", q, k) * self.scale + if mask is not None: + scores = jnp.where(mask[:, None, None, :], scores, -1e9) + if self.causal: + causal_mask = jnp.tril(jnp.ones((seq_len, seq_len), dtype=bool)) + scores = jnp.where(causal_mask, scores, -1e9) + wmask = self._window_mask(seq_len) + if wmask is not None: + scores = jnp.where(wmask, scores, -1e9) + weights = jax.nn.softmax(scores, axis=-1) + context = jnp.einsum("bhqk,bhkd->bhqd", weights, v) + context = jnp.swapaxes(context, 1, 2).reshape(batch, seq_len, self.embed_dim) + out = self.out_proj(context) + if mask is not None: + out = out * mask[..., None] + + out = shard(out, self.shd_cfg.act_btd) + return out + + +class TransformerLayer(nnx.Module): + def __init__(self, d_model: int, attention_heads: int, ffn_dim: int, causal: bool, + attn_window_size: Tuple[int, int], shd_cfg: MiMoShardingCfg, dtype=jnp.float32, + rngs: Optional[nnx.Rngs] = None): + self.act = jax.nn.gelu + self.shd_cfg = shd_cfg + + self.self_attn = Attention(d_model, attention_heads, attn_window_size, causal, + shd_cfg, dtype=dtype, rngs=rngs) + + self.self_attn_layer_norm = shard( + nnx.LayerNorm(d_model, epsilon=1e-6, param_dtype=dtype, rngs=rngs), + shd_cfg.norm_scale + ) + self.final_layer_norm = shard( + nnx.LayerNorm(d_model, epsilon=1e-6, param_dtype=dtype, rngs=rngs), + shd_cfg.norm_scale + ) + + self.fc1 = shard( + nnx.Linear(d_model, ffn_dim, dtype=dtype, rngs=rngs), + shd_cfg.ffn_weight_in + ) + self.fc2 = shard( + nnx.Linear(ffn_dim, d_model, dtype=dtype, rngs=rngs), + shd_cfg.ffn_weight_out + ) + + def __call__(self, hidden_states: Array, mask: Optional[Array], rope: Optional[Tuple[Array, Array]]) -> Array: + residual = hidden_states + hidden_states = self.self_attn_layer_norm(hidden_states) + hidden_states = self.self_attn(hidden_states, mask, rope) + hidden_states = residual + hidden_states + residual = hidden_states + hidden_states = self.final_layer_norm(hidden_states) + hidden_states = self.fc1(hidden_states) + hidden_states = self.act(hidden_states) + hidden_states = shard(hidden_states, self.shd_cfg.act_btd) + hidden_states = self.fc2(hidden_states) + return residual + hidden_states + + +class ResidualVectorQuantizer(nnx.Module): + def __init__(self, dimension: int, n_q: int, bins: Sequence[int], + shd_cfg: MiMoShardingCfg, dtype=jnp.float32, + rngs: Optional[nnx.Rngs] = None): + self.dimension = dimension + self.n_q = n_q + self.shd_cfg = shd_cfg + + codebooks_list = [] + for i in range(n_q): + size = bins[min(i, len(bins) - 1)] + embed = jnp.zeros((size, dimension), dtype=dtype) + codebooks_list.append(shard(nnx.Param(embed), shd_cfg.codebook)) + self.codebooks = nnx.List(codebooks_list) + + def encode(self, hidden_states: Array, mask: Optional[Array] = None, n_q: Optional[int] = None) -> Tuple[ + Array, Array]: + num_levels = n_q or self.n_q + residual = hidden_states + quantized = jnp.zeros_like(hidden_states) + codes = [] + mask = None if mask is None else mask[..., None] + for i in range(num_levels): + codebook = self.codebooks[i].value + dist = jnp.sum((residual[:, None, :] - codebook[None, :, :]) ** 2, axis=-1) + idx = jnp.argmin(dist, axis=-1) + chosen = codebook[idx] + if mask is not None: + chosen = chosen * mask + quantized = quantized + chosen + residual = residual - chosen + codes.append(idx) + return jnp.stack(codes, axis=0), quantized + + def decode(self, codes: Array) -> Array: + num_levels = codes.shape[0] + flat = codes.reshape(num_levels, -1) + decoded = jnp.zeros((flat.shape[1], self.dimension), dtype=jnp.float32) + for i in range(num_levels): + codebook = self.codebooks[i].value + decoded = decoded + codebook[flat[i]] + return decoded.reshape(*codes.shape[1:], self.dimension) + + +class AudioEncoder(nnx.Module): + def __init__(self, config: MiMoAudioTokenizerConfig, dtype=jnp.float32, rngs: Optional[nnx.Rngs] = None): + self.config = config + self.shd_cfg = config.shd_cfg + + self.conv1 = shard( + nnx.Conv( + in_features=config.n_mels, + out_features=config.d_model, + kernel_size=config.kernel_size, + padding="SAME", + param_dtype=dtype, + rngs=rngs + ), + self.shd_cfg.conv_weight + ) + self.conv2 = shard( + nnx.Conv( + in_features=config.d_model, + out_features=config.d_model, + kernel_size=config.kernel_size, + strides=config.stride_size, + padding="SAME", + param_dtype=dtype, + rngs=rngs + ), + self.shd_cfg.conv_weight + ) + + self.position_embedding = RotaryEmbedding(config.rope_theta, config.d_model // config.encoder_attention_heads, + config.max_audio_seconds * config.sampling_rate // config.hop_length, + config.rope_type, dtype=dtype) + + self.layers = nnx.List([ + TransformerLayer(config.d_model, config.encoder_attention_heads, config.encoder_ffn_dim, + config.encoder_causal, tuple(config.encoder_attn_window_size), + self.shd_cfg, dtype=dtype, rngs=rngs) + for _ in range(config.encoder_layers) + ]) + + self.layer_norm = shard( + nnx.LayerNorm(config.d_model, epsilon=1e-6, param_dtype=dtype, rngs=rngs), + self.shd_cfg.norm_scale + ) + + if config.avg_pooler != 1: + self.down_sample_layer = shard( + nnx.Conv( + in_features=config.d_model, + out_features=config.d_model, + kernel_size=config.avg_pooler, + strides=config.avg_pooler, + padding="SAME", + use_bias=False, + param_dtype=dtype, + rngs=rngs + ), + self.shd_cfg.conv_weight + ) + self.down_norm = shard( + nnx.LayerNorm(config.d_model, epsilon=1e-6, param_dtype=dtype, rngs=rngs), + self.shd_cfg.norm_scale + ) + else: + self.down_sample_layer = None + self.down_norm = None + + if config.num_quantizers: + bins = config.codebook_size or [1024] + self.quantizer = ResidualVectorQuantizer(config.d_model, config.num_quantizers, bins, + self.shd_cfg, dtype=dtype, rngs=rngs) + else: + self.quantizer = None + + def get_output_length(self, mel_len: Array) -> Array: + tgt = mel_len + 3 - self.config.kernel_size + return (tgt + 2 - self.config.kernel_size) // self.config.stride_size + 1 + + def __call__(self, input_features: Array, input_lens: Array, use_quantizer: bool = True, + n_q: Optional[int] = None) -> EncoderOutput: + x = input_features + x = jax.nn.gelu(self.conv1(x)) + x = shard(x, self.shd_cfg.act_btd) + + x = jax.nn.gelu(self.conv2(x)) + x = shard(x, self.shd_cfg.act_btd) + + lengths = self.get_output_length(input_lens) + max_len = x.shape[1] + mask = make_sequence_mask(lengths, max_len) + pos = get_position_ids(lengths, max_len) + rope = self.position_embedding(x, pos) + skip = 0.0 + for idx, layer in enumerate(self.layers): + x = layer(x, mask, rope) + if self.config.encoder_skip_layer_id and idx == self.config.encoder_skip_layer_id - 1: + skip = x + x = x + skip + x = self.layer_norm(x) + if self.down_sample_layer is not None: + x = jax.nn.gelu(self.down_sample_layer(x)) + x = shard(x, self.shd_cfg.act_btd) + + lengths = (lengths // self.config.avg_pooler) + ((lengths % self.config.avg_pooler) != 0).astype( + lengths.dtype) + max_len = x.shape[1] + mask = make_sequence_mask(lengths, max_len) + x = self.down_norm(x) + x = x * mask[..., None] + packed = x.reshape(-1, self.config.d_model) + mask_flat = mask.reshape(-1) + codes = None + if self.quantizer is not None and use_quantizer: + codes, quantized = self.quantizer.encode(packed, mask=mask_flat, n_q=n_q) + packed = quantized + packed = packed.reshape(x.shape) + return EncoderOutput(hidden_states=packed, packed_states=packed, output_lengths=lengths, codes=codes) + + def decode_vq(self, codes: Array) -> Array: + if self.quantizer is None: + raise ValueError("Quantizer disabled") + return self.quantizer.decode(codes) + + +class CausalConvTranspose1d(nnx.Module): + def __init__(self, in_channels: int, out_channels: int, kernel_size: int, stride: int, + shd_cfg: MiMoShardingCfg | None = None, + dtype=jnp.float32, rngs: Optional[nnx.Rngs] = None): + self.shd_cfg = shd_cfg or MiMoShardingCfg.no_sharding() + + self.conv = ConvTranspose1d(in_channels, out_channels, kernel_size, stride, + shd_cfg=self.shd_cfg, dtype=dtype, rngs=rngs) + + self.norm = shard( + nnx.GroupNorm(num_features=out_channels, num_groups=1, epsilon=1e-5, + param_dtype=dtype, rngs=rngs), + self.shd_cfg.norm_scale + ) + + self.kernel_size = kernel_size + self.stride = stride + + def __call__(self, x: Array, input_length: Array) -> Tuple[Array, Array]: + y = self.conv(x) + y = self.norm(y) + trim = max(0, self.kernel_size - self.stride) + if trim > 0: + y = y[:, :-trim, :] + output_len = (input_length - 1) * self.stride + self.kernel_size - trim + return y, output_len + + +class TransformerVocos(nnx.Module): + def __init__(self, config: MiMoAudioTokenizerConfig, dtype=jnp.float32, rngs: Optional[nnx.Rngs] = None): + self.config = config + self.shd_cfg = config.shd_cfg + + self.embeddings = shard( + nnx.Linear(config.n_mels, config.vocoder_dim, use_bias=False, dtype=dtype, rngs=rngs), + self.shd_cfg.attn_qkvo_weight + ) + + self.position_embedding = RotaryEmbedding(config.rope_theta, + config.vocoder_dim // config.vocoder_attention_heads, + config.max_audio_seconds * config.sampling_rate // config.hop_length, + config.rope_type, dtype=dtype) + + self.layers = nnx.List([ + TransformerLayer(config.vocoder_dim, config.vocoder_attention_heads, config.vocoder_intermediate_dim, + False, tuple(config.vocoder_attn_window_size), + self.shd_cfg, dtype=dtype, rngs=rngs) + for _ in range(config.vocoder_num_layers) + ]) + + self.layer_norm = shard( + nnx.LayerNorm(config.vocoder_dim, epsilon=1e-6, param_dtype=dtype, rngs=rngs), + self.shd_cfg.norm_scale + ) + + self.head = ISTFTHead(config.vocoder_dim, config.nfft, config.hop_length, + config.vocoder_padding, shd_cfg=self.shd_cfg, + dtype=dtype, rngs=rngs) + + def __call__(self, mels: Array, input_length: Array) -> VocoderOutput: + x = self.embeddings(mels) + mask = make_sequence_mask(input_length, x.shape[1]) + pos = get_position_ids(input_length, x.shape[1]) + rope = self.position_embedding(x, pos) + for layer in self.layers: + x = layer(x, mask, rope) + x = self.layer_norm(x) + x = x * mask[..., None] + wav = self.head(x) + wav_len = input_length * self.config.hop_length + wav = wav[:, None, :] + return VocoderOutput(wav=wav, wav_lengths=wav_len) + + +class AudioDecoder(nnx.Module): + def __init__(self, config: MiMoAudioTokenizerConfig, dtype=jnp.float32, rngs: Optional[nnx.Rngs] = None): + self.config = config + self.shd_cfg = config.shd_cfg + + if config.avg_pooler != 1: + self.dconv1 = CausalConvTranspose1d(config.d_model, config.d_model, config.avg_pooler, + config.avg_pooler, shd_cfg=self.shd_cfg, + dtype=dtype, rngs=rngs) + else: + self.dconv1 = None + + self.position_embedding = RotaryEmbedding(config.rope_theta, config.d_model // config.decoder_attention_heads, + config.max_audio_seconds * config.sampling_rate // config.hop_length, + config.rope_type, dtype=dtype) + + self.layers = nnx.List([ + TransformerLayer(config.d_model, config.decoder_attention_heads, config.decoder_ffn_dim, + config.decoder_causal, tuple(config.decoder_attn_window_size), + self.shd_cfg, dtype=dtype, rngs=rngs) + for _ in range(config.decoder_layers) + ]) + + self.layer_norm = shard( + nnx.LayerNorm(config.d_model, epsilon=1e-6, param_dtype=dtype, rngs=rngs), + self.shd_cfg.norm_scale + ) + + self.dconv2 = CausalConvTranspose1d(config.d_model, config.n_mels, config.decoder_kernel_size, + config.decoder_stride_size, shd_cfg=self.shd_cfg, + dtype=dtype, rngs=rngs) + + self.vocoder = TransformerVocos(config, dtype=dtype, rngs=rngs) + + def __call__(self, audio_embed: Array, input_length: Array) -> Array: + x = audio_embed + lengths = input_length + if self.dconv1 is not None: + x, lengths = self.dconv1(x, lengths) + mask = make_sequence_mask(lengths, x.shape[1]) + pos = get_position_ids(lengths, x.shape[1]) + rope = self.position_embedding(x, pos) + for layer in self.layers: + x = layer(x, mask, rope) + x = self.layer_norm(x) + coarse, mel_lengths = self.dconv2(x, lengths) + vocoder_out = self.vocoder(coarse, mel_lengths) + return vocoder_out.wav + + +class FlaxMiMoAudioTokenizer(nnx.Module): + def __init__(self, config: MiMoAudioTokenizerConfig, dtype=jnp.float32, rngs: Optional[nnx.Rngs] = None): + self.config = config + self.encoder = AudioEncoder(config, dtype=dtype, rngs=rngs) + self.decoder = AudioDecoder(config, dtype=dtype, rngs=rngs) + self.downsample_rate = int(config.hop_length * 2 * config.avg_pooler) + + def __call__(self, mels: Array, input_lens: Array, use_quantizer: bool = True) -> Array: + enc = self.encoder(mels, input_lens, use_quantizer=use_quantizer) + return self.decoder(enc.hidden_states, enc.output_lengths) + + def encode(self, mels: Array, input_lens: Array, use_quantizer: bool = True, + n_q: Optional[int] = None) -> EncoderOutput: + return self.encoder(mels, input_lens, use_quantizer=use_quantizer, n_q=n_q) + + def decode(self, codes: Array) -> Array: + hidden = self.encoder.decode_vq(codes) + hidden = hidden[None, ...] + lengths = jnp.array([hidden.shape[1]]) + return self.decoder(hidden, lengths) diff --git a/bonsai/models/mimo_audio/mimo_audio_tokenizer_configuration.py b/bonsai/models/mimo_audio/mimo_audio_tokenizer_configuration.py new file mode 100644 index 00000000..2512fc27 --- /dev/null +++ b/bonsai/models/mimo_audio/mimo_audio_tokenizer_configuration.py @@ -0,0 +1,227 @@ +"""Configuration classes for MiMo Audio Tokenizer.""" + +from dataclasses import dataclass +from typing import Optional +import jax.numpy as jnp +from transformers import PretrainedConfig +from jax.sharding import PartitionSpec as P + +Array = jnp.ndarray +ShardingSpec = P + + +@dataclass(slots=True, frozen=True) +class MiMoShardingCfg: + """Sharding configuration for MiMo Audio Tokenizer. + + Controls how model parameters and activations are distributed across devices. + """ + # Conv layer weight sharding + conv_weight: ShardingSpec # (in_channels, out_channels, kernel_size) + conv_bias: ShardingSpec # (out_channels,) + + # Transformer weight sharding (shared by Encoder/Decoder/Vocoder) + attn_qkvo_weight: ShardingSpec # (d_model, d_model) + attn_qkv_bias: ShardingSpec # (d_model,) + attn_out_bias: ShardingSpec # (d_model,) + + # FFN weight sharding + ffn_weight_in: ShardingSpec # (d_model, ffn_dim) + ffn_weight_out: ShardingSpec # (ffn_dim, d_model) + ffn_bias: ShardingSpec # (ffn_dim,) or (d_model,) + + # LayerNorm/GroupNorm sharding + norm_scale: ShardingSpec # (dim,) + norm_bias: ShardingSpec # (dim,) + + # Quantizer codebook sharding + codebook: ShardingSpec # (codebook_size, d_model) + + # ConvTranspose1d weight sharding + conv_transpose_weight: ShardingSpec # (in_ch, out_ch, kernel) + conv_transpose_bias: ShardingSpec # (out_ch,) + + # ISTFT related sharding + istft_linear_weight: ShardingSpec # (dim, n_fft+2) + istft_linear_bias: ShardingSpec # (n_fft+2,) + istft_window: ShardingSpec # (win_length,) + + # Activation sharding + act_btd: ShardingSpec # [batch, time, d_model] + act_btnh: ShardingSpec # [batch, time, num_heads, head_dim] + act_btc: ShardingSpec # [batch, time, channels] + + @staticmethod + def no_sharding(): + """Configuration with no sharding (all None).""" + return MiMoShardingCfg( + conv_weight=P(None, None, None), + conv_bias=P(None), + attn_qkvo_weight=P(None, None), + attn_qkv_bias=P(None), + attn_out_bias=P(None), + ffn_weight_in=P(None, None), + ffn_weight_out=P(None, None), + ffn_bias=P(None), + norm_scale=P(None), + norm_bias=P(None), + codebook=P(None, None), + conv_transpose_weight=P(None, None, None), + conv_transpose_bias=P(None), + istft_linear_weight=P(None, None), + istft_linear_bias=P(None), + istft_window=P(None), + act_btd=P(None, None, None), + act_btnh=P(None, None, None, None), + act_btc=P(None, None, None), + ) + + @staticmethod + def default(): + """Default sharding configuration for distributed training.""" + return MiMoShardingCfg( + conv_weight=P(None, "tp", None), + conv_bias=P("tp"), + attn_qkvo_weight=P("fsdp", "tp"), + attn_qkv_bias=P("tp"), + attn_out_bias=P("tp"), + ffn_weight_in=P("fsdp", "tp"), + ffn_weight_out=P("tp", "fsdp"), + ffn_bias=P("tp"), + norm_scale=P("tp"), + norm_bias=P("tp"), + codebook=P("tp", "fsdp"), + conv_transpose_weight=P(None, "tp", None), + conv_transpose_bias=P("tp"), + istft_linear_weight=P("fsdp", "tp"), + istft_linear_bias=P("tp"), + istft_window=P(None), # replicated + act_btd=P("fsdp", None, "tp"), + act_btnh=P("fsdp", None, "tp", None), + act_btc=P("fsdp", None, "tp"), + ) + + +class MiMoAudioTokenizerConfig(PretrainedConfig): + model_type = "mimo_audio_tokenizer" + + def __init__( + self, + max_audio_seconds: int = 1800, + stride_size: int = 2, + avg_pooler: int = 2, + d_model: int = 1280, + scale_embedding: bool = False, + kernel_size: int = 3, + activation_function: str = "gelu", + encoder_layers: int = 32, + encoder_skip_layer_id: int = 3, + encoder_attention_heads: int = 20, + encoder_ffn_dim: int = 5120, + encoder_causal: bool = False, + encoder_attn_window_size: list[int] = None, # [-1,-1] + decoder_layers: int = 32, + decoder_attention_heads: int = 20, + decoder_ffn_dim: int = 5120, + decoder_kernel_size: int = 3, + decoder_stride_size: int = 2, + decoder_causal: bool = True, + decoder_attn_window_size: list[int] = None, # [-1,-1] + nfft: int = 960, + vocoder_dim: int = 256, + vocoder_intermediate_dim: int = 1024, + vocoder_num_layers: int = 16, + n_mels: int = 128, + sampling_rate: int = 24000, + hop_length: int = 240, + window_size: int = 960, + vocoder_padding: str = "same", + fmin: int = 0, + fmax: int = None, + num_quantizers: int = 20, + codebook_size: list[int] = None, + # [1024,1024,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128] + threshold_ema_dead_code: int = 2, + position_embedding_type: str = "rope", + rope_theta: int = 10000, + rope_type: str = "default", + ln_type: str = "LayerNorm", + vocoder_attention_heads: int = 16, + vocoder_attn_window_size: list[int] = None, # [40,10] + use_sharding: bool = False, + shd_cfg: MiMoShardingCfg | None = None, + **kwargs, + ): + super().__init__(**kwargs) + self.max_audio_seconds = max_audio_seconds + self.stride_size = stride_size + self.avg_pooler = avg_pooler + self.d_model = d_model + self.scale_embedding = scale_embedding + self.kernel_size = kernel_size + self.activation_function = activation_function + self.encoder_layers = encoder_layers + self.encoder_skip_layer_id = encoder_skip_layer_id + self.encoder_attention_heads = encoder_attention_heads + self.encoder_ffn_dim = encoder_ffn_dim + self.encoder_causal = encoder_causal + self.encoder_attn_window_size = ( + encoder_attn_window_size + if encoder_attn_window_size is not None + else [-1, -1] + ) + self.decoder_layers = decoder_layers + self.decoder_attention_heads = decoder_attention_heads + self.decoder_ffn_dim = decoder_ffn_dim + self.decoder_kernel_size = decoder_kernel_size + self.decoder_stride_size = decoder_stride_size + self.decoder_causal = decoder_causal + self.decoder_attn_window_size = ( + decoder_attn_window_size + if decoder_attn_window_size is not None + else [-1, -1] + ) + self.nfft = nfft + self.vocoder_dim = vocoder_dim + self.vocoder_intermediate_dim = vocoder_intermediate_dim + self.vocoder_num_layers = vocoder_num_layers + self.n_mels = n_mels + self.sampling_rate = sampling_rate + self.hop_length = hop_length + self.window_size = window_size + self.vocoder_padding = vocoder_padding + self.fmin = fmin + self.fmax = fmax + self.num_quantizers = num_quantizers + self.codebook_size = codebook_size if codebook_size is not None else [1024] + self.threshold_ema_dead_code = threshold_ema_dead_code + self.position_embedding_type = position_embedding_type + self.rope_theta = rope_theta + self.rope_type = rope_type + self.ln_type = ln_type + self.vocoder_attention_heads = vocoder_attention_heads + self.vocoder_attn_window_size = ( + vocoder_attn_window_size + if vocoder_attn_window_size is not None + else [40, 10] + ) + + # Sharding configuration + if shd_cfg is None: + self.shd_cfg = MiMoShardingCfg.default() if use_sharding else MiMoShardingCfg.no_sharding() + else: + self.shd_cfg = shd_cfg + + +@dataclass +class EncoderOutput: + hidden_states: Array + packed_states: Array + output_lengths: Array + codes: Optional[Array] + + +@dataclass +class VocoderOutput: + wav: Array + wav_lengths: Array diff --git a/bonsai/models/mimo_audio/mimo_audio_tokenizer_params.py b/bonsai/models/mimo_audio/mimo_audio_tokenizer_params.py new file mode 100644 index 00000000..9035da73 --- /dev/null +++ b/bonsai/models/mimo_audio/mimo_audio_tokenizer_params.py @@ -0,0 +1,249 @@ +"""Weight loading for MiMo Audio Tokenizer from SafeTensors""" + +import gc +import re +from dataclasses import dataclass +from typing import Any + +import jax +import jax.numpy as jnp +from flax import nnx +from safetensors import safe_open + +from bonsai.models.mimo_audio import mimo_audio_tokenizer as model_lib + + +@dataclass(frozen=True) +class Transform: + permute: tuple[int, ...] | None = None + reshape: tuple[int, ...] | None = None + reshape_first: bool = False + + +TRANSFORM_LINEAR = Transform(permute=(1, 0)) +TRANSFORM_CONV1D = Transform(permute=(2, 1, 0)) +TRANSFORM_NONE = Transform() + + +def _get_key_mapping(config: model_lib.MiMoAudioTokenizerConfig) -> dict[str, tuple[str, Transform]]: + mapping = { + r"encoder\.conv1\.weight": ("encoder.conv1.kernel", TRANSFORM_CONV1D), + r"encoder\.conv1\.bias": ("encoder.conv1.bias", TRANSFORM_NONE), + r"encoder\.conv2\.weight": ("encoder.conv2.kernel", TRANSFORM_CONV1D), + r"encoder\.conv2\.bias": ("encoder.conv2.bias", TRANSFORM_NONE), + r"encoder\.layer_norm\.weight": ("encoder.layer_norm.scale", TRANSFORM_NONE), + r"encoder\.layer_norm\.bias": ("encoder.layer_norm.bias", TRANSFORM_NONE), + } + + for idx in range(config.encoder_layers): + layer_mappings = { + rf"encoder\.layers\.{idx}\.self_attn\.q_proj\.weight": (f"encoder.layers.{idx}.self_attn.q_proj.kernel", TRANSFORM_LINEAR), + rf"encoder\.layers\.{idx}\.self_attn\.q_proj\.bias": (f"encoder.layers.{idx}.self_attn.q_proj.bias", TRANSFORM_NONE), + rf"encoder\.layers\.{idx}\.self_attn\.k_proj\.weight": (f"encoder.layers.{idx}.self_attn.k_proj.kernel", TRANSFORM_LINEAR), + rf"encoder\.layers\.{idx}\.self_attn\.k_proj\.bias": (f"encoder.layers.{idx}.self_attn.k_proj.bias", TRANSFORM_NONE), + rf"encoder\.layers\.{idx}\.self_attn\.v_proj\.weight": (f"encoder.layers.{idx}.self_attn.v_proj.kernel", TRANSFORM_LINEAR), + rf"encoder\.layers\.{idx}\.self_attn\.v_proj\.bias": (f"encoder.layers.{idx}.self_attn.v_proj.bias", TRANSFORM_NONE), + rf"encoder\.layers\.{idx}\.self_attn\.out_proj\.weight": (f"encoder.layers.{idx}.self_attn.out_proj.kernel", TRANSFORM_LINEAR), + rf"encoder\.layers\.{idx}\.self_attn\.out_proj\.bias": (f"encoder.layers.{idx}.self_attn.out_proj.bias", TRANSFORM_NONE), + rf"encoder\.layers\.{idx}\.self_attn_layer_norm\.weight": (f"encoder.layers.{idx}.self_attn_layer_norm.scale", TRANSFORM_NONE), + rf"encoder\.layers\.{idx}\.self_attn_layer_norm\.bias": (f"encoder.layers.{idx}.self_attn_layer_norm.bias", TRANSFORM_NONE), + rf"encoder\.layers\.{idx}\.final_layer_norm\.weight": (f"encoder.layers.{idx}.final_layer_norm.scale", TRANSFORM_NONE), + rf"encoder\.layers\.{idx}\.final_layer_norm\.bias": (f"encoder.layers.{idx}.final_layer_norm.bias", TRANSFORM_NONE), + rf"encoder\.layers\.{idx}\.fc1\.weight": (f"encoder.layers.{idx}.fc1.kernel", TRANSFORM_LINEAR), + rf"encoder\.layers\.{idx}\.fc1\.bias": (f"encoder.layers.{idx}.fc1.bias", TRANSFORM_NONE), + rf"encoder\.layers\.{idx}\.fc2\.weight": (f"encoder.layers.{idx}.fc2.kernel", TRANSFORM_LINEAR), + rf"encoder\.layers\.{idx}\.fc2\.bias": (f"encoder.layers.{idx}.fc2.bias", TRANSFORM_NONE), + } + mapping.update(layer_mappings) + + mapping.update({ + r"encoder\.down_sample_layer\.0\.weight": ("encoder.down_sample_layer.kernel", TRANSFORM_CONV1D), + r"encoder\.down_sample_layer\.0\.bias": ("encoder.down_sample_layer.bias", TRANSFORM_NONE), + r"encoder\.down_sample_norm\.weight": ("encoder.down_norm.scale", TRANSFORM_NONE), + r"encoder\.down_sample_norm\.bias": ("encoder.down_norm.bias", TRANSFORM_NONE), + }) + + for idx in range(config.num_quantizers): + mapping[rf"encoder\.quantizer\.vq\.layers\.{idx}\._codebook\.embed"] = ( + f"encoder.quantizer.codebooks.{idx}", + TRANSFORM_NONE, + ) + + mapping.update({ + r"decoder\.dconv1\.conv\.weight": ("decoder.dconv1.conv.kernel", TRANSFORM_NONE), + r"decoder\.dconv1\.conv\.bias": ("decoder.dconv1.conv.bias", TRANSFORM_NONE), + r"decoder\.dconv1\.norm\.weight": ("decoder.dconv1.norm.scale", TRANSFORM_NONE), + r"decoder\.dconv1\.norm\.bias": ("decoder.dconv1.norm.bias", TRANSFORM_NONE), + r"decoder\.layer_norm\.weight": ("decoder.layer_norm.scale", TRANSFORM_NONE), + r"decoder\.layer_norm\.bias": ("decoder.layer_norm.bias", TRANSFORM_NONE), + r"decoder\.dconv2\.conv\.weight": ("decoder.dconv2.conv.kernel", TRANSFORM_NONE), + r"decoder\.dconv2\.conv\.bias": ("decoder.dconv2.conv.bias", TRANSFORM_NONE), + r"decoder\.dconv2\.norm\.weight": ("decoder.dconv2.norm.scale", TRANSFORM_NONE), + r"decoder\.dconv2\.norm\.bias": ("decoder.dconv2.norm.bias", TRANSFORM_NONE), + }) + + for idx in range(config.decoder_layers): + layer_mappings = { + rf"decoder\.layers\.{idx}\.self_attn\.q_proj\.weight": (f"decoder.layers.{idx}.self_attn.q_proj.kernel", TRANSFORM_LINEAR), + rf"decoder\.layers\.{idx}\.self_attn\.q_proj\.bias": (f"decoder.layers.{idx}.self_attn.q_proj.bias", TRANSFORM_NONE), + rf"decoder\.layers\.{idx}\.self_attn\.k_proj\.weight": (f"decoder.layers.{idx}.self_attn.k_proj.kernel", TRANSFORM_LINEAR), + rf"decoder\.layers\.{idx}\.self_attn\.k_proj\.bias": (f"decoder.layers.{idx}.self_attn.k_proj.bias", TRANSFORM_NONE), + rf"decoder\.layers\.{idx}\.self_attn\.v_proj\.weight": (f"decoder.layers.{idx}.self_attn.v_proj.kernel", TRANSFORM_LINEAR), + rf"decoder\.layers\.{idx}\.self_attn\.v_proj\.bias": (f"decoder.layers.{idx}.self_attn.v_proj.bias", TRANSFORM_NONE), + rf"decoder\.layers\.{idx}\.self_attn\.out_proj\.weight": (f"decoder.layers.{idx}.self_attn.out_proj.kernel", TRANSFORM_LINEAR), + rf"decoder\.layers\.{idx}\.self_attn\.out_proj\.bias": (f"decoder.layers.{idx}.self_attn.out_proj.bias", TRANSFORM_NONE), + rf"decoder\.layers\.{idx}\.self_attn_layer_norm\.weight": (f"decoder.layers.{idx}.self_attn_layer_norm.scale", TRANSFORM_NONE), + rf"decoder\.layers\.{idx}\.self_attn_layer_norm\.bias": (f"decoder.layers.{idx}.self_attn_layer_norm.bias", TRANSFORM_NONE), + rf"decoder\.layers\.{idx}\.final_layer_norm\.weight": (f"decoder.layers.{idx}.final_layer_norm.scale", TRANSFORM_NONE), + rf"decoder\.layers\.{idx}\.final_layer_norm\.bias": (f"decoder.layers.{idx}.final_layer_norm.bias", TRANSFORM_NONE), + rf"decoder\.layers\.{idx}\.fc1\.weight": (f"decoder.layers.{idx}.fc1.kernel", TRANSFORM_LINEAR), + rf"decoder\.layers\.{idx}\.fc1\.bias": (f"decoder.layers.{idx}.fc1.bias", TRANSFORM_NONE), + rf"decoder\.layers\.{idx}\.fc2\.weight": (f"decoder.layers.{idx}.fc2.kernel", TRANSFORM_LINEAR), + rf"decoder\.layers\.{idx}\.fc2\.bias": (f"decoder.layers.{idx}.fc2.bias", TRANSFORM_NONE), + } + mapping.update(layer_mappings) + + mapping.update({ + r"decoder\.vocoder\.embeddings\.weight": ("decoder.vocoder.embeddings.kernel", TRANSFORM_LINEAR), + r"decoder\.vocoder\.embeddings\.bias": ("decoder.vocoder.embeddings.bias", TRANSFORM_NONE), + r"decoder\.vocoder\.layer_norm\.weight": ("decoder.vocoder.layer_norm.scale", TRANSFORM_NONE), + r"decoder\.vocoder\.layer_norm\.bias": ("decoder.vocoder.layer_norm.bias", TRANSFORM_NONE), + r"decoder\.vocoder\.head\.out\.weight": ("decoder.vocoder.head.linear.kernel", TRANSFORM_LINEAR), + r"decoder\.vocoder\.head\.out\.bias": ("decoder.vocoder.head.linear.bias", TRANSFORM_NONE), + r"decoder\.vocoder\.head\.istft\.window": ("decoder.vocoder.head.istft.window", TRANSFORM_NONE), + }) + + for idx in range(config.vocoder_num_layers): + layer_mappings = { + rf"decoder\.vocoder\.layers\.{idx}\.self_attn\.q_proj\.weight": (f"decoder.vocoder.layers.{idx}.self_attn.q_proj.kernel", TRANSFORM_LINEAR), + rf"decoder\.vocoder\.layers\.{idx}\.self_attn\.q_proj\.bias": (f"decoder.vocoder.layers.{idx}.self_attn.q_proj.bias", TRANSFORM_NONE), + rf"decoder\.vocoder\.layers\.{idx}\.self_attn\.k_proj\.weight": (f"decoder.vocoder.layers.{idx}.self_attn.k_proj.kernel", TRANSFORM_LINEAR), + rf"decoder\.vocoder\.layers\.{idx}\.self_attn\.k_proj\.bias": (f"decoder.vocoder.layers.{idx}.self_attn.k_proj.bias", TRANSFORM_NONE), + rf"decoder\.vocoder\.layers\.{idx}\.self_attn\.v_proj\.weight": (f"decoder.vocoder.layers.{idx}.self_attn.v_proj.kernel", TRANSFORM_LINEAR), + rf"decoder\.vocoder\.layers\.{idx}\.self_attn\.v_proj\.bias": (f"decoder.vocoder.layers.{idx}.self_attn.v_proj.bias", TRANSFORM_NONE), + rf"decoder\.vocoder\.layers\.{idx}\.self_attn\.out_proj\.weight": (f"decoder.vocoder.layers.{idx}.self_attn.out_proj.kernel", TRANSFORM_LINEAR), + rf"decoder\.vocoder\.layers\.{idx}\.self_attn\.out_proj\.bias": (f"decoder.vocoder.layers.{idx}.self_attn.out_proj.bias", TRANSFORM_NONE), + rf"decoder\.vocoder\.layers\.{idx}\.self_attn_layer_norm\.weight": (f"decoder.vocoder.layers.{idx}.self_attn_layer_norm.scale", TRANSFORM_NONE), + rf"decoder\.vocoder\.layers\.{idx}\.self_attn_layer_norm\.bias": (f"decoder.vocoder.layers.{idx}.self_attn_layer_norm.bias", TRANSFORM_NONE), + rf"decoder\.vocoder\.layers\.{idx}\.final_layer_norm\.weight": (f"decoder.vocoder.layers.{idx}.final_layer_norm.scale", TRANSFORM_NONE), + rf"decoder\.vocoder\.layers\.{idx}\.final_layer_norm\.bias": (f"decoder.vocoder.layers.{idx}.final_layer_norm.bias", TRANSFORM_NONE), + rf"decoder\.vocoder\.layers\.{idx}\.fc1\.weight": (f"decoder.vocoder.layers.{idx}.fc1.kernel", TRANSFORM_LINEAR), + rf"decoder\.vocoder\.layers\.{idx}\.fc1\.bias": (f"decoder.vocoder.layers.{idx}.fc1.bias", TRANSFORM_NONE), + rf"decoder\.vocoder\.layers\.{idx}\.fc2\.weight": (f"decoder.vocoder.layers.{idx}.fc2.kernel", TRANSFORM_LINEAR), + rf"decoder\.vocoder\.layers\.{idx}\.fc2\.bias": (f"decoder.vocoder.layers.{idx}.fc2.bias", TRANSFORM_NONE), + } + mapping.update(layer_mappings) + + return mapping + + +def _get_jax_key( + mapping: dict[str, tuple[str, Transform]], source_key: str +) -> tuple[str | None, Transform | None]: + for pat, (jax_key, transform) in mapping.items(): + if re.fullmatch(pat, source_key): + return jax_key, transform + return None, None + + +def _assign_weights( + keys: list[str | int], + tensor: Any, + state_dict: dict, + transform: Transform | None, + sharding_dict: dict | None = None, +) -> None: + key, *rest = keys + if not rest: + if transform is not None: + if transform.reshape_first and transform.reshape is not None: + tensor = tensor.reshape(transform.reshape) + if transform.permute is not None: + tensor = tensor.transpose(transform.permute) + if not transform.reshape_first and transform.reshape is not None: + tensor = tensor.reshape(transform.reshape) + + if tensor.shape != state_dict[key].shape: + raise ValueError(f"Shape mismatch: {tensor.shape} vs {state_dict[key].shape}") + + target_dtype = state_dict[key].dtype + if sharding_dict is not None: + state_dict[key] = jax.device_put(jnp.asarray(tensor, dtype=target_dtype), sharding_dict[key]) + else: + state_dict[key] = jax.device_put(jnp.asarray(tensor, dtype=target_dtype)) + else: + next_sharding = sharding_dict[key] if sharding_dict is not None else None + _assign_weights(rest, tensor, state_dict[key], transform, next_sharding) + + +def _stoi(s: str) -> str | int: + try: + return int(s) + except ValueError: + return s + + +def load_tokenizer_weights_from_safetensors( + config: model_lib.MiMoAudioTokenizerConfig, + safetensors_path: str, + dtype=jnp.float32, + mesh: jax.sharding.Mesh | None = None, + rngs: nnx.Rngs | None = None, +) -> model_lib.FlaxMiMoAudioTokenizer: + + model = nnx.eval_shape(lambda: model_lib.FlaxMiMoAudioTokenizer(config, dtype=dtype, rngs=nnx.Rngs(params=0))) + graph_def, abs_state = nnx.split(model) + state_dict = abs_state.to_pure_dict() + + sharding = nnx.get_named_sharding(abs_state, mesh).to_pure_dict() if mesh is not None else None + + key_mapping = _get_key_mapping(config) + conversion_errors = [] + skipped_keys = [] + + with safe_open(safetensors_path, framework="numpy") as sf: + for torch_key in sf.keys(): + jax_key, transform = _get_jax_key(key_mapping, torch_key) + if jax_key is None: + skipped_keys.append(torch_key) + continue + + keys = [_stoi(k) for k in jax_key.split(".")] + try: + tensor = sf.get_tensor(torch_key) + _assign_weights(keys, tensor, state_dict, transform, sharding) + except Exception as e: + conversion_errors.append(f"Failed to assign '{torch_key}' to '{jax_key}': {type(e).__name__}: {e}") + + if skipped_keys: + print(f"Warning: Skipped {len(skipped_keys)} keys not found in mapping") + + if conversion_errors: + raise RuntimeError( + f"Encountered {len(conversion_errors)} weight conversion errors:\n" + "\n".join(conversion_errors) + ) + + model = nnx.merge(graph_def, state_dict) + + encoder_rotary_dim = config.d_model // config.encoder_attention_heads + encoder_half_dim = encoder_rotary_dim // 2 + encoder_inv_freq = 1.0 / (config.rope_theta ** (jnp.arange(0, encoder_half_dim, dtype=jnp.float32) / float(encoder_half_dim))) + model.encoder.position_embedding.inv_freq.value = encoder_inv_freq + + decoder_rotary_dim = config.d_model // config.decoder_attention_heads + decoder_half_dim = decoder_rotary_dim // 2 + decoder_inv_freq = 1.0 / (config.rope_theta ** (jnp.arange(0, decoder_half_dim, dtype=jnp.float32) / float(decoder_half_dim))) + model.decoder.position_embedding.inv_freq.value = decoder_inv_freq + + vocoder_rotary_dim = config.vocoder_dim // config.vocoder_attention_heads + vocoder_half_dim = vocoder_rotary_dim // 2 + vocoder_inv_freq = 1.0 / (config.rope_theta ** (jnp.arange(0, vocoder_half_dim, dtype=jnp.float32) / float(vocoder_half_dim))) + model.decoder.vocoder.position_embedding.inv_freq.value = vocoder_inv_freq + + window = jnp.hanning(config.nfft).astype(dtype) + model.decoder.vocoder.head.istft.window.value = window + + gc.collect() + + print("Tokenizer weights loaded successfully!") + return model From a64b513ad5c8c83913dc4f2de2acbe868b44f585 Mon Sep 17 00:00:00 2001 From: haibo Date: Mon, 19 Jan 2026 13:02:02 +0800 Subject: [PATCH 06/18] add mimo-audio readme --- bonsai/models/mimo_audio/README.md | 94 ++++++++++++++++++++++++++++++ 1 file changed, 94 insertions(+) create mode 100644 bonsai/models/mimo_audio/README.md diff --git a/bonsai/models/mimo_audio/README.md b/bonsai/models/mimo_audio/README.md new file mode 100644 index 00000000..2c5f7893 --- /dev/null +++ b/bonsai/models/mimo_audio/README.md @@ -0,0 +1,94 @@ +# MiMo-Audio in JAX + +This directory contains a pure JAX implementation of the [MiMo-Audio multimodal language model](https://github.com/XiaomiMiMo/MiMo), using the [Flax NNX](https://flax.readthedocs.io/en/v0.8.3/experimental/nnx/index.html) API. + +MiMo-Audio is a unified speech-text model that supports: +- **Text-to-Speech (TTS)**: Generate natural speech from text +- **Speech-to-Text (ASR)**: Transcribe speech to text +- **Speech-to-Speech**: Direct speech translation and conversion + +## Model Configuration Support Status + +| Model Name | Config Support Status | +| :--- | :--- | +| **Main Models** | | +| [MiMo-Audio-7B-Base](https://huggingface.co/XiaomiMiMo/MiMo-Audio-7B-Base) | **✅ Supported** | +| [MiMo-Audio-7B-Instruct](https://huggingface.co/XiaomiMiMo/MiMo-Audio-7B-Instruct) | **✅ Supported** | +| **Audio Tokenizer** | | +| [MiMo-Audio-Tokenizer](https://huggingface.co/XiaomiMiMo/MiMo-Audio-Tokenizer) | **✅ Supported** | + + +## Running this model + +Run MiMo-Audio inference with a minimal example: + +```sh +python3 -m bonsai.models.mimo_audio.test.run_model +``` +## Model Architecture + +MiMo-Audio consists of three main components: + +### 1. Main Transformer (Qwen2-based) +- **Layers**: 36 +- **Hidden size**: 4096 +- **Attention heads**: 32 (8 KV heads) +- **Intermediate size**: 11008 +- **Max position embeddings**: 8192 + +### 2. Local Transformer (for audio generation) +- **Layers**: 16 +- **Hidden size**: 1024 +- **Attention heads**: 64 +- **FFN dimension**: 4096 + +### 3. Input Local Transformer (for audio encoding) +- **Layers**: 6 +- **Hidden size**: 1024 +- **Attention heads**: 64 +- **Bidirectional attention**: Yes + +### 4. Audio Tokenizer +- **Encoder layers**: 32 +- **Decoder layers**: 32 +- **Quantizers**: 20 +- **Sampling rate**: 24000 Hz +- **Audio channels**: 8 (used for multi-codebook representation) + +## Special Tokens + +MiMo-Audio uses special tokens for controlling speech generation: + +- `<|sostm|>` (151648): Start of speech/stream +- `<|eostm|>` (151649): End of speech/stream +- `<|sosp|>` (151646): Start of speech +- `<|eosp|>` (151647): End of speech +- `<|empty|>` (151645): Empty token (indicates audio generation) +- `<|eot|>` (151643): End of turn + +## Input Format + +MiMo-Audio uses an interleaved format where each group contains: +- 1 text token (repeated `group_size` times) +- 8 audio channel tokens (one per channel, repeated `group_size` times) + +Shape: `[batch, audio_channels + 1, num_groups * group_size]` + +For text-only input (TTS), audio channels are filled with channel-specific empty IDs. + +## Audio Processing + +### Encoding (Speech → Tokens) +1. Waveform → Mel spectrogram +2. Mel spectrogram → Encoder +3. Encoder → Quantizer → Audio tokens (8 channels) + +### Decoding (Tokens → Speech) +1. Audio tokens → Decoder +2. Decoder → Vocoder +3. Vocoder → Waveform (24kHz) + +## References + +- [MiMo-Audio Paper](https://github.com/XiaomiMiMo/MiMo-Audio/blob/main/MiMo-Audio-Technical-Report.pdf) +- [Official Implementation](https://github.com/XiaomiMiMo/MiMo-Audio) From f0cc7a0361db232b7f618a820250f6464195fba0 Mon Sep 17 00:00:00 2001 From: haibo Date: Mon, 19 Jan 2026 13:02:25 +0800 Subject: [PATCH 07/18] test: mimo-audio --- bonsai/models/mimo_audio/test/__init__.py | 0 bonsai/models/mimo_audio/test/run_model.py | 248 +++++++++++++++++++++ 2 files changed, 248 insertions(+) create mode 100644 bonsai/models/mimo_audio/test/__init__.py create mode 100644 bonsai/models/mimo_audio/test/run_model.py diff --git a/bonsai/models/mimo_audio/test/__init__.py b/bonsai/models/mimo_audio/test/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/bonsai/models/mimo_audio/test/run_model.py b/bonsai/models/mimo_audio/test/run_model.py new file mode 100644 index 00000000..9df13176 --- /dev/null +++ b/bonsai/models/mimo_audio/test/run_model.py @@ -0,0 +1,248 @@ +#!/usr/bin/env python3 + +import os +import json +import jax +import jax.numpy as jnp +import numpy as np +from flax import nnx + + +def load_audio_tokenizer(tokenizer_path: str): + from bonsai.models.mimo_audio.mimo_audio_tokenizer_configuration import MiMoAudioTokenizerConfig + from bonsai.models.mimo_audio.mimo_audio_tokenizer_params import load_tokenizer_weights_from_safetensors + + config_path = os.path.join(tokenizer_path, "config.json") + with open(config_path) as f: + config_dict = json.load(f) + + config_dict['use_sharding'] = False + config = MiMoAudioTokenizerConfig(**config_dict) + safetensors_path = os.path.join(tokenizer_path, "model.safetensors") + + tokenizer_model = load_tokenizer_weights_from_safetensors( + config=config, + safetensors_path=safetensors_path, + dtype=jnp.float32, + mesh=None, + rngs=nnx.Rngs(0), + ) + + return tokenizer_model, config + +def load_main_model(model_path: str): + from bonsai.models.mimo_audio.mimo_audio_configuration import MiMoAudioConfig, MiMoAudioArguments + from bonsai.models.mimo_audio.params import create_model_with_weights + from transformers import AutoTokenizer + + config_path = os.path.join(model_path, "config.json") + with open(config_path) as f: + config_dict = json.load(f) + config_kwargs = {k: v for k, v in config_dict.items() if k in MiMoAudioConfig.__dataclass_fields__} + config = MiMoAudioConfig(**config_kwargs) + + text_tokenizer = AutoTokenizer.from_pretrained(model_path) + + args = MiMoAudioArguments( + model_name_or_path=model_path, + sosp_idx=text_tokenizer.convert_tokens_to_ids("<|sosp|>"), + eosp_idx=text_tokenizer.convert_tokens_to_ids("<|eosp|>"), + sostm_idx=text_tokenizer.convert_tokens_to_ids("<|sostm|>"), + eostm_idx=text_tokenizer.convert_tokens_to_ids("<|eostm|>"), + eot_idx=text_tokenizer.convert_tokens_to_ids("<|eot|>"), + empty_idx=text_tokenizer.convert_tokens_to_ids("<|empty|>"), + ) + model = create_model_with_weights( + model_path=model_path, + config=config, + args=args, + rngs=nnx.Rngs(0), + mesh=None, + ) + + return model, config, args, text_tokenizer + +def insert_between(tokens: list, group_size: int, fill_value: int) -> list: + if group_size <= 1: + return tokens + + result = [] + for token in tokens: + result.append(token) + result.extend([fill_value] * (group_size - 1)) + + return result + + +def run_inference( + main_model, + tokenizer_model, + text_tokenizer, + config, + args, + tokenizer_config, + text_to_speak: str, + max_steps: int = 100, + output_dir: str = "test_outputs" +): + from bonsai.models.mimo_audio.modeling import forward_jit, MiMoSampler + from bonsai.models.mimo_audio.mimo_audio_configuration import MiMoSamplerConfig + + audio_channels = main_model.audio_channels + group_size = main_model.group_size + batch_size = 1 + + tts_template = "Turn this writing into audio" + chat_text = f"<|im_start|>user\n{tts_template}: {text_to_speak}<|im_end|>\n<|im_start|>assistant\n<|sostm|>" + + text_tokens_raw = text_tokenizer.encode(chat_text) + text_tokens_with_spacing = insert_between(text_tokens_raw, group_size, -100) + + num_groups = len(text_tokens_with_spacing) // group_size + if len(text_tokens_with_spacing) % group_size != 0: + text_tokens_with_spacing.extend([-100] * (group_size - len(text_tokens_with_spacing) % group_size)) + num_groups = len(text_tokens_with_spacing) // group_size + + input_shape = (batch_size, audio_channels + 1, num_groups * group_size) + input_ids = jnp.zeros(input_shape, dtype=jnp.int32) + + input_ids = input_ids.at[0, 0, :].set(jnp.array(text_tokens_with_spacing)) + + for ch in range(1, audio_channels + 1): + channel_empty_id = main_model.speech_empty_ids[ch - 1] + audio_empty_tokens = jnp.full((num_groups * group_size,), channel_empty_id, dtype=jnp.int32) + input_ids = input_ids.at[0, ch, :].set(audio_empty_tokens) + + cache = main_model.model.init_cache( + main_model.qwen2_config, + batch_size, + num_groups, + generate_steps=max_steps, + dtype=jnp.bfloat16, + ) + + text_sampler = MiMoSampler(MiMoSamplerConfig(temperature=0.6, top_p=1.0, do_sample=True)) + audio_sampler = MiMoSampler(MiMoSamplerConfig(temperature=0.9, top_p=0.95, do_sample=True)) + + + pad_id = text_tokenizer.pad_token_id + text_logits, local_hidden_states, cache = forward_jit( + main_model, input_ids, cache, pad_id + ) + + generated_text_tokens = [] + generated_audio_tokens_list = [] + + rng_key = jax.random.key(42) + empty_idx = args.empty_idx + + for step in range(max_steps): + key, subkey = jax.random.split(rng_key) + logits_2d = text_logits[0, 0:1, :] + next_text_token = text_sampler.sample(logits_2d, subkey) + next_text_token_int = int(next_text_token[0]) + generated_text_tokens.append(next_text_token_int) + + if next_text_token_int == args.eostm_idx: + break + if next_text_token_int == text_tokenizer.eos_token_id: + break + + audio_tokens = None + + if next_text_token_int != empty_idx: + for t in range(group_size): + audio_tokens_step = jnp.array(main_model.speech_empty_ids) + generated_audio_tokens_list.append(audio_tokens_step) + else: + key, subkey = jax.random.split(key) + audio_tokens = main_model.local_forward( + local_hidden_states, + subkey, + audio_sampler + ) + + for t in range(group_size): + audio_tokens_step = audio_tokens[0, t, :] + generated_audio_tokens_list.append(audio_tokens_step) + + rng_key = key + + next_input = jnp.zeros((batch_size, audio_channels + 1, group_size), dtype=jnp.int32) + + for i in range(group_size): + next_input = next_input.at[0, 0, i].set(next_text_token[0]) + + if audio_tokens is None: + for ch in range(audio_channels): + channel_empty_id = main_model.speech_empty_ids[ch] + for i in range(group_size): + next_input = next_input.at[0, ch + 1, i].set(channel_empty_id) + else: + for ch in range(audio_channels): + for i in range(group_size): + next_input = next_input.at[0, ch + 1, i].set(audio_tokens[0, i, ch]) + + text_logits, local_hidden_states, cache = forward_jit( + main_model, next_input, cache, pad_id + ) + + generated_text = text_tokenizer.decode(generated_text_tokens, skip_special_tokens=True) + print(f"text token output: {generated_text}") + + + audio_tokens_array = jnp.stack(generated_audio_tokens_list, axis=0).T + + speech_empty_ids = main_model.speech_empty_ids + is_real_audio_mask = jnp.zeros(audio_tokens_array.shape[1], dtype=bool) + + for ch in range(audio_channels): + empty_id = speech_empty_ids[ch] + not_empty = audio_tokens_array[ch, :] != empty_id + is_real_audio_mask = is_real_audio_mask | not_empty + + + audio_tokens_array = audio_tokens_array[:, is_real_audio_mask] + decoded_audio = tokenizer_model.decode(audio_tokens_array) + os.makedirs(output_dir, exist_ok=True) + + import soundfile as sf + audio_path = os.path.join(output_dir, "generated_audio.wav") + audio_np = np.array(decoded_audio[0, 0, :]) + sample_rate = tokenizer_config.sampling_rate + sf.write(audio_path, audio_np, sample_rate) + + print(f"\n wav file saved: {audio_path}") + + +def main(): + + """Replace it with the real model file location""" + model_path = os.path.expanduser( + "~/.cache/modelscope/hub/models/XiaomiMiMo/MiMo-Audio-7B-Instruct" + ) + tokenizer_path = os.path.expanduser( + "~/.cache/modelscope/hub/models/XiaomiMiMo/MiMo-Audio-Tokenizer" + ) + + tokenizer_model, tokenizer_config = load_audio_tokenizer(tokenizer_path) + main_model, config, args, text_tokenizer = load_main_model(model_path) + + text_to_speak = ("And now here is my secret, a very simple secret:It is only with the heart that one can see rightly;" + "What is essential is invisible to the eye.It's the time you wasted for your rose that makes your rose so important." + "Men have forgotten this truth, but you must not forget it.You become responsible for what you have tamed." + "You are responsible for your rose...") + run_inference( + main_model=main_model, + tokenizer_model=tokenizer_model, + text_tokenizer=text_tokenizer, + config=config, + args=args, + tokenizer_config=tokenizer_config, + text_to_speak=text_to_speak, + max_steps=300, + ) + + +if __name__ == "__main__": + main() From f67af1e1b4106ddc2a0fca3fa4bd21d3c30c052d Mon Sep 17 00:00:00 2001 From: haibo Date: Mon, 19 Jan 2026 14:36:16 +0800 Subject: [PATCH 08/18] test: qwen2 --- .../models/qwen2/tests/test_outputs_qwen2.py | 368 ++++++++++++++++++ 1 file changed, 368 insertions(+) create mode 100644 bonsai/models/qwen2/tests/test_outputs_qwen2.py diff --git a/bonsai/models/qwen2/tests/test_outputs_qwen2.py b/bonsai/models/qwen2/tests/test_outputs_qwen2.py new file mode 100644 index 00000000..219fbe5c --- /dev/null +++ b/bonsai/models/qwen2/tests/test_outputs_qwen2.py @@ -0,0 +1,368 @@ +import jax +import jax.numpy as jnp +import numpy as np +import torch +from absl.testing import absltest +from flax import nnx +from huggingface_hub import snapshot_download +from jax.sharding import AxisType +from jax.typing import DTypeLike +from transformers import AutoTokenizer +from transformers.cache_utils import DynamicCache +from transformers.models.qwen2 import Qwen2ForCausalLM + +from bonsai.models.qwen2 import modeling, params +from bonsai.models.qwen2.tests.run_model import tokenize + + +class TestModuleForwardPasses(absltest.TestCase): + def setUp(self): + super().setUp() + jax.config.update("jax_default_matmul_precision", "float32") + model_name: str = "Qwen/Qwen2-0.5B" + self.tokenizer = AutoTokenizer.from_pretrained(model_name) + + self.torch_model = Qwen2ForCausalLM.from_pretrained(model_name, torch_dtype=torch.float32).eval() + self.bonsai_config = modeling.ModelConfig.qwen2_0_5b(use_sharding=False) + model_ckpt_path = snapshot_download("Qwen/Qwen2-0.5B") + self.mesh = jax.make_mesh(((1, 1)), ("fsdp", "tp"), axis_types=(AxisType.Explicit, AxisType.Explicit)) + jax.set_mesh(self.mesh) + + graph_def, state = nnx.split( + params.create_model_from_safe_tensors(model_ckpt_path, self.bonsai_config, self.mesh) + ) + state = jax.tree.map(lambda x: x.astype(jnp.float32) if isinstance(x, jax.Array) else x, state) + self.nnx_model = nnx.merge(graph_def, state) + + self.batch_size = 32 + self.num_input_tokens = 5 + self.cache_size, self.gen_steps = 128, 10 + self.relaxed_tol = 1e-3 + + def _check_batched_logits(self, left_pads: int, torch_logits: torch.Tensor, nnx_logits: jax.Array): + max_len = torch_logits.shape[-2] + for lp, tl, nl in zip(left_pads, torch_logits, nnx_logits): + torch.testing.assert_close( + torch.tensor(np.array(nl, dtype=np.float32))[lp:max_len, :], + tl[lp:, :], + rtol=self.relaxed_tol, + atol=self.relaxed_tol, + check_dtype=False, + ) + + def _setup_torch_attn(self, input_embeddings: torch.Tensor, attention_mask: None = None): + past_key_values = DynamicCache() + past_seen_tokens = past_key_values.get_seq_length() if past_key_values is not None else 0 + cache_position = torch.arange( + past_seen_tokens, past_seen_tokens + input_embeddings.shape[1], device=self.torch_model.device + ) + position_ids = cache_position.unsqueeze(0) + + batch_size, seq_length = input_embeddings.shape[:2] + if attention_mask is None: + attention_mask = torch.ones((batch_size, seq_length), dtype=torch.bool, device=input_embeddings.device) + + causal_mask = torch.triu( + torch.ones((seq_length, seq_length), dtype=torch.bool, device=input_embeddings.device), + diagonal=1 + ) + causal_mask = causal_mask.unsqueeze(0).unsqueeze(0) + causal_mask = causal_mask.expand(batch_size, 1, seq_length, seq_length) + causal_mask = torch.where(causal_mask, torch.finfo(input_embeddings.dtype).min, 0.0) + + position_embeddings = self.torch_model.model.rotary_emb(input_embeddings, position_ids) + out = dict( + hidden_states=input_embeddings.to(torch.float32), + attention_mask=causal_mask, + position_ids=position_ids, + past_key_values=past_key_values, + use_cache=True, + cache_position=cache_position, + position_embeddings=position_embeddings, + ) + return out + + def _nnx_forward_logits(self, cache: modeling.Cache, tokens: jax.Array, dtype: DTypeLike = jnp.float32): + segment_ids = 1 * (tokens != self.tokenizer.pad_token_id) + x = self.nnx_model.embedder.embedding.value.at[(tokens,)].get().astype(jnp.float32) + for i, layer in enumerate(self.nnx_model.layers): + x = layer(x, cache[i], segment_ids).astype(dtype) + nnx_logits = self.nnx_model.lm_head(self.nnx_model.final_norm(x)) + return nnx_logits + + def _process_hf_tokens(self, query: list[str]): + messages = [{"role": "user", "content": s} for s in query] + text = [ + self.tokenizer.apply_chat_template([m], tokenize=False, add_generation_prompt=True) for m in messages + ] + model_inputs = self.tokenizer(text, return_tensors="pt", padding=True, padding_side="left").to( + self.torch_model.device + ) + tmp = model_inputs["attention_mask"] + num_zeros = tmp.shape[1] - tmp.sum(dim=-1) + model_inputs["left_pads"] = num_zeros + pos_ids = torch.arange(tmp.shape[1]) - num_zeros.reshape(-1, 1) + pos_ids[pos_ids < 0] = 2**30 + model_inputs["position_ids"] = pos_ids + return model_inputs + + def _init_nnx_cache(self, batch_size: int): + return self.nnx_model.init_cache( + cfg=self.bonsai_config, batch_size=batch_size, token_len=10, generate_steps=32, dtype=jnp.float32 + ) + + def test_embedder(self): + nm = self.nnx_model.embedder + tm = self.torch_model.model.embed_tokens + + tx = torch.randint(0, self.torch_model.config.vocab_size, size=(self.batch_size, self.num_input_tokens)) + jx = jnp.array(tx.cpu().detach().numpy()) + + jy, ty = nm.embedding.value.at[(jx,)].get(), tm(tx) + torch.testing.assert_close( + torch.tensor(np.array(jy, dtype=np.float32)), + ty, + rtol=self.relaxed_tol, + atol=self.relaxed_tol, + check_dtype=False, + ) + + def test_decoder_layer(self): + nm = self.nnx_model.layers[0] + tm = self.torch_model.model.layers[0].to(torch.float32) + + shape = (self.batch_size, self.num_input_tokens, self.bonsai_config.emb_dim) + jx = jax.random.normal(jax.random.key(0), shape=shape) + tx = torch.tensor(np.array(jx, dtype=np.float32)) + nnx_cache = self._init_nnx_cache(self.batch_size) + torch_inputs = self._setup_torch_attn(tx) + + jy, ty = nm(jx, nnx_cache[0], jnp.ones((self.batch_size, self.num_input_tokens))), tm(**torch_inputs) + torch.testing.assert_close( + torch.tensor(np.array(jy, dtype=np.float32)), + ty, + rtol=self.relaxed_tol, + atol=self.relaxed_tol, + check_dtype=False, + ) + + def test_all_decoder_layers(self): + nnx_cache = self._init_nnx_cache(self.batch_size) + shape = (self.batch_size, self.num_input_tokens, self.bonsai_config.emb_dim) + + for nm, tm, nc in zip(self.nnx_model.layers, self.torch_model.model.layers, nnx_cache): + jx = jax.random.normal(jax.random.key(0), shape=shape) + tx = torch.tensor(np.array(jx, dtype=np.float32)) + + jy = nm(jx, nc, jnp.ones((self.batch_size, self.num_input_tokens))) + torch_inputs = self._setup_torch_attn(tx) + ty = tm.to(torch.float32)(**torch_inputs) + torch.testing.assert_close( + torch.tensor(np.array(jy, dtype=np.float32)), + ty, + atol=self.relaxed_tol, + rtol=self.relaxed_tol, + check_dtype=False, + ) + + def test_rms_norm(self): + nm = self.nnx_model.layers[0].input_layernorm + tm = self.torch_model.model.layers[0].input_layernorm + + shape = (self.batch_size, self.num_input_tokens, self.bonsai_config.emb_dim) + jx = jax.random.normal(jax.random.key(0), shape=shape) + tx = torch.tensor(np.array(jx, dtype=np.float32)) + + jy, ty = nm(jx), tm(tx) + torch.testing.assert_close( + torch.tensor(np.array(jy, dtype=np.float32)), + ty, + rtol=self.relaxed_tol, + atol=self.relaxed_tol, + check_dtype=False, + ) + + def test_self_attn(self): + nm = self.nnx_model.layers[0].attn + tm = self.torch_model.model.layers[0].self_attn.to(torch.float32) + + shape = (self.batch_size, self.num_input_tokens, self.bonsai_config.emb_dim) + jx = jax.random.normal(jax.random.key(0), shape=shape) + tx = torch.tensor(np.array(jx, dtype=np.float32)) + torch_inputs = self._setup_torch_attn(tx) + nnx_cache = self._init_nnx_cache(self.batch_size) + + jy = nm(jx, nnx_cache[0], jnp.ones((self.batch_size, self.num_input_tokens), dtype=jnp.float32)) + ty = tm(**torch_inputs)[0] + torch.testing.assert_close( + torch.tensor(np.array(jy, dtype=np.float32)), + ty, + rtol=self.relaxed_tol, + atol=self.relaxed_tol, + check_dtype=False, + ) + + def test_q_proj(self): + nm = self.nnx_model.layers[0].attn.q_proj + tm = self.torch_model.model.layers[0].self_attn.q_proj.to(torch.float32) + + shape = (self.batch_size, self.num_input_tokens, self.bonsai_config.emb_dim) + jx = jax.random.normal(jax.random.key(0), shape=shape) + tx = torch.tensor(np.array(jx, dtype=np.float32)) + + jy = nm(jx) + ty = tm(tx) + + jy_flat = jy.reshape(self.batch_size, self.num_input_tokens, -1) + torch.testing.assert_close( + torch.tensor(np.array(jy_flat, dtype=np.float32)), + ty, + rtol=self.relaxed_tol, + atol=self.relaxed_tol, + check_dtype=False, + ) + + def test_k_proj(self): + nm = self.nnx_model.layers[0].attn.k_proj + tm = self.torch_model.model.layers[0].self_attn.k_proj.to(torch.float32) + + shape = (self.batch_size, self.num_input_tokens, self.bonsai_config.emb_dim) + jx = jax.random.normal(jax.random.key(0), shape=shape) + tx = torch.tensor(np.array(jx, dtype=np.float32)) + + jy = nm(jx) + ty = tm(tx) + + jy_flat = jy.reshape(self.batch_size, self.num_input_tokens, -1) + torch.testing.assert_close( + torch.tensor(np.array(jy_flat, dtype=np.float32)), + ty, + rtol=self.relaxed_tol, + atol=self.relaxed_tol, + check_dtype=False, + ) + + def test_o_proj(self): + nm = self.nnx_model.layers[0].attn.o_proj + tm = self.torch_model.model.layers[0].self_attn.o_proj.to(torch.float32) + + input_dim = self.bonsai_config.num_heads * self.bonsai_config.head_dim + shape = (self.batch_size, self.num_input_tokens, input_dim) + + jx = jax.random.normal(jax.random.key(0), shape=shape) + tx = torch.tensor(np.array(jx, dtype=np.float32)) + + jy = nm(jx) + ty = tm(tx) + torch.testing.assert_close( + torch.tensor(np.array(jy, dtype=np.float32)), + ty, + rtol=self.relaxed_tol, + atol=self.relaxed_tol, + check_dtype=False, + ) + + def test_mlp(self): + nm = self.nnx_model.layers[0].mlp + tm = self.torch_model.model.layers[0].mlp.to(torch.float32) + + shape = (self.batch_size, self.num_input_tokens, self.bonsai_config.emb_dim) + jx = jax.random.normal(jax.random.key(0), shape=shape) + tx = torch.tensor(np.array(jx, dtype=np.float32)) + + jy, ty = nm(jx), tm(tx) + torch.testing.assert_close( + torch.tensor(np.array(jy, dtype=np.float32)), + ty, + rtol=self.relaxed_tol, + atol=self.relaxed_tol, + check_dtype=False, + ) + + def test_lm_head(self): + nm = self.nnx_model.lm_head + tm = self.torch_model.lm_head.to(torch.float32) + + shape = (self.batch_size, self.num_input_tokens, self.bonsai_config.emb_dim) + jx = jax.random.normal(jax.random.key(0), shape=shape) + tx = torch.tensor(np.array(jx, dtype=np.float32)) + + jy, ty = nm(jx), tm(tx) + torch.testing.assert_close( + torch.tensor(np.array(jy, dtype=np.float32)), + ty, + rtol=self.relaxed_tol, + atol=self.relaxed_tol, + check_dtype=False, + ) + + def test_sin_cos(self): + batch_size, seq_len = 2, 10 + dim = self.bonsai_config.head_dim + hidden_states = torch.ones((batch_size, seq_len, dim)) + jp = jnp.stack([jnp.arange(seq_len), jnp.arange(seq_len)]) + js, jc = modeling._generate_pos_embeddings(jp, dim, self.bonsai_config.rope_theta) + + position_ids = torch.tensor(np.array(jp, dtype=np.int64)) + torch_rope_output = self.torch_model.model.rotary_emb(hidden_states, position_ids) + + if isinstance(torch_rope_output, tuple) and len(torch_rope_output) == 2: + tc, ts = torch_rope_output + else: + tc = torch_rope_output + ts = torch_rope_output + + half_dim = dim // 2 + if tc.shape[-1] != half_dim: + tc = tc[:, :, :half_dim] + if ts.shape[-1] != half_dim: + ts = ts[:, :, :half_dim] + + torch.testing.assert_close( + torch.tensor(np.array(js, dtype=np.float32)), + ts, + rtol=self.relaxed_tol, + atol=self.relaxed_tol, + check_dtype=False, + ) + torch.testing.assert_close( + torch.tensor(np.array(jc, dtype=np.float32)), + tc, + rtol=self.relaxed_tol, + atol=self.relaxed_tol, + check_dtype=False, + ) + + def test_full(self): + query = ["Why is the sky blue instead of any other color like purple?"] + tokens = tokenize(self.tokenizer, query) + _, token_len = tokens.shape + self.torch_model = self.torch_model.to(torch.float32) + nnx_cache = self._init_nnx_cache(len(query)) + + nnx_logits = self._nnx_forward_logits(nnx_cache, tokens, jnp.float32) + torch_inputs = self._process_hf_tokens(query) + torch_logits = self.torch_model(**torch_inputs).logits + torch.testing.assert_close( + torch.tensor(np.array(nnx_logits, dtype=np.float32))[:, :token_len, :], + torch_logits, + rtol=self.relaxed_tol, + atol=self.relaxed_tol, + check_dtype=False, + ) + + def test_full_batched(self): + query = ["Why is the sky blue instead of any other color like purple?", "Who am I?"] + tokens = tokenize(self.tokenizer, query) + self.torch_model = self.torch_model.to(torch.float32) + nnx_cache = self._init_nnx_cache(len(query)) + + nnx_logits = self._nnx_forward_logits(nnx_cache, tokens, jnp.float32) + torch_inputs = self._process_hf_tokens(query) + torch_logits = self.torch_model(**torch_inputs).logits + + self._check_batched_logits(torch_inputs["left_pads"], torch_logits, nnx_logits) + + +if __name__ == "__main__": + absltest.main() From ff21293770252c09b0c1b1ebe81218ad61ad7da0 Mon Sep 17 00:00:00 2001 From: haibo Date: Mon, 19 Jan 2026 19:46:09 +0800 Subject: [PATCH 09/18] test: mimo-audio --- bonsai/models/mimo_audio/test/run_model.py | 13 +- .../test/test_outputs_mimo_audio.py | 344 ++++++++++++++++++ bonsai/models/qwen2/tests/run_model.py | 9 +- 3 files changed, 354 insertions(+), 12 deletions(-) create mode 100644 bonsai/models/mimo_audio/test/test_outputs_mimo_audio.py diff --git a/bonsai/models/mimo_audio/test/run_model.py b/bonsai/models/mimo_audio/test/run_model.py index 9df13176..6ac7e215 100644 --- a/bonsai/models/mimo_audio/test/run_model.py +++ b/bonsai/models/mimo_audio/test/run_model.py @@ -6,6 +6,7 @@ import jax.numpy as jnp import numpy as np from flax import nnx +from huggingface_hub import snapshot_download def load_audio_tokenizer(tokenizer_path: str): @@ -217,13 +218,11 @@ def run_inference( def main(): - """Replace it with the real model file location""" - model_path = os.path.expanduser( - "~/.cache/modelscope/hub/models/XiaomiMiMo/MiMo-Audio-7B-Instruct" - ) - tokenizer_path = os.path.expanduser( - "~/.cache/modelscope/hub/models/XiaomiMiMo/MiMo-Audio-Tokenizer" - ) + model_name = "XiaomiMiMo/MiMo-Audio-7B-Instruct" + tokenizer_name = "XiaomiMiMo/MiMo-Audio-Tokenizer" + + model_path = snapshot_download(model_name) + tokenizer_path = snapshot_download(tokenizer_name) tokenizer_model, tokenizer_config = load_audio_tokenizer(tokenizer_path) main_model, config, args, text_tokenizer = load_main_model(model_path) diff --git a/bonsai/models/mimo_audio/test/test_outputs_mimo_audio.py b/bonsai/models/mimo_audio/test/test_outputs_mimo_audio.py new file mode 100644 index 00000000..95001350 --- /dev/null +++ b/bonsai/models/mimo_audio/test/test_outputs_mimo_audio.py @@ -0,0 +1,344 @@ +import json +import os +from dataclasses import asdict + +import jax +import jax.numpy as jnp +import numpy as np +import torch +from absl.testing import absltest +from flax import nnx +from huggingface_hub import snapshot_download +from transformers import AutoTokenizer +from transformers.cache_utils import DynamicCache +from transformers.masking_utils import create_causal_mask + +from bonsai.models.mimo_audio import params +from bonsai.models.mimo_audio.mimo_audio_configuration import MiMoAudioConfig, MiMoAudioArguments +from bonsai.models.mimo_audio.pytorch.src.mimo_audio.modeling_mimo_audio import ( + MiMoAudioForCausalLM as TorchMiMoAudio, + MiMoAudioConfig as TorchMiMoAudioConfig, +) +from bonsai.models.qwen3.modeling import ShardingCfg + + +class TestMiMoAudioLayerOutputs(absltest.TestCase): + @classmethod + def setUpClass(cls): + super().setUpClass() + jax.config.update("jax_default_matmul_precision", "float32") + jax.config.update("jax_platforms", "cpu") + + model_name = "XiaomiMiMo/MiMo-Audio-7B-Instruct" + cls.tokenizer = AutoTokenizer.from_pretrained(model_name) + model_ckpt_path = snapshot_download(model_name) + + config_path = os.path.join(model_ckpt_path, "config.json") + with open(config_path) as f: + config_dict = json.load(f) + + config_kwargs = {k: v for k, v in config_dict.items() if k in MiMoAudioConfig.__dataclass_fields__} + config_kwargs['shd_cfg'] = ShardingCfg.no_sharding() + cls.bonsai_config = MiMoAudioConfig(**config_kwargs) + + cls.args = MiMoAudioArguments( + model_name_or_path=model_ckpt_path, + sosp_idx=cls.tokenizer.convert_tokens_to_ids("<|sosp|>"), + eosp_idx=cls.tokenizer.convert_tokens_to_ids("<|eosp|>"), + sostm_idx=cls.tokenizer.convert_tokens_to_ids("<|sostm|>"), + eostm_idx=cls.tokenizer.convert_tokens_to_ids("<|eostm|>"), + eot_idx=cls.tokenizer.convert_tokens_to_ids("<|eot|>"), + empty_idx=cls.tokenizer.convert_tokens_to_ids("<|empty|>"), + ) + + torch_config = TorchMiMoAudioConfig.from_pretrained(model_ckpt_path) + cls.torch_model = TorchMiMoAudio.from_pretrained( + model_ckpt_path, config=torch_config, args=asdict(cls.args), torch_dtype=torch.float32 + ).eval().cpu() + + cls.nnx_model = params.create_model_with_weights( + model_path=model_ckpt_path, config=cls.bonsai_config, args=cls.args, + rngs=nnx.Rngs(0), dtype=jnp.float32, mesh=None, + ) + + cls.batch_size = 1 + cls.num_input_tokens = 5 + cls.group_size = cls.bonsai_config.group_size + cls.audio_channels = cls.bonsai_config.audio_channels + cls.tol = 1e-3 + + def _init_cache(self, batch_size, token_len): + return self.nnx_model.model.init_cache( + cfg=self.nnx_model.qwen2_config, batch_size=batch_size, + token_len=token_len, generate_steps=0, dtype=jnp.float32 + ) + + def _compare(self, jy, ty): + if ty.dim() == 2 and jy.ndim == 3: + ty = ty.unsqueeze(0) + torch.testing.assert_close( + torch.tensor(np.array(jy, dtype=np.float32)), ty, + rtol=self.tol, atol=self.tol, check_dtype=False, + ) + + def test_text_embedder(self): + tx = torch.randint(0, self.torch_model.config.vocab_size, size=(self.batch_size, self.num_input_tokens)) + jx = jnp.array(tx.cpu().detach().numpy()) + jy = self.nnx_model.model.embedder.embedding.value[jx] + with torch.no_grad(): + ty = self.torch_model.model.embed_tokens(tx) + self._compare(jy, ty) + + def test_speech_embeddings(self): + for ch in range(self.audio_channels): + vocab_size = self.nnx_model.speech_vocab_sizes[ch] + tx = torch.randint(0, vocab_size, size=(self.batch_size, self.num_input_tokens)) + jx = jnp.array(tx.cpu().detach().numpy()) + jy = self.nnx_model.speech_embeddings[ch](jx) + with torch.no_grad(): + ty = self.torch_model.speech_embeddings[ch](tx) + self._compare(jy, ty) + + def test_main_decoder_layer(self): + shape = (self.batch_size, self.num_input_tokens, self.bonsai_config.hidden_size) + jx = jax.random.normal(jax.random.key(0), shape=shape) + tx = torch.tensor(np.array(jx, dtype=np.float32)) + + cache = self._init_cache(self.batch_size, self.num_input_tokens) + segment_ids = jnp.ones((self.batch_size, self.num_input_tokens)) + jy = self.nnx_model.model.layers[0](jx, cache[0], segment_ids) + + cache_position = torch.arange(0, self.num_input_tokens, device=tx.device) + position_ids = cache_position.unsqueeze(0) + position_embeddings = self.torch_model.model.rotary_emb(tx, position_ids) + + with torch.no_grad(): + ty_output = self.torch_model.model.layers[0].to(torch.float32)( + tx, position_ids=position_ids, position_embeddings=position_embeddings, + past_key_value=DynamicCache(), cache_position=cache_position, + ) + ty = ty_output[0] if isinstance(ty_output, tuple) else ty_output + + self._compare(jy, ty) + + def test_all_main_decoder_layers(self): + cache = self._init_cache(self.batch_size, self.num_input_tokens) + shape = (self.batch_size, self.num_input_tokens, self.bonsai_config.hidden_size) + + for layer_idx, (nm, tm, nc) in enumerate(zip( + self.nnx_model.model.layers, self.torch_model.model.layers, cache + )): + jx = jax.random.normal(jax.random.key(layer_idx), shape=shape) + tx = torch.tensor(np.array(jx, dtype=np.float32)) + segment_ids = jnp.ones((self.batch_size, self.num_input_tokens)) + jy = nm(jx, nc, segment_ids) + + cache_position = torch.arange(0, self.num_input_tokens, device=tx.device) + position_ids = cache_position.unsqueeze(0) + position_embeddings = self.torch_model.model.rotary_emb(tx, position_ids) + + with torch.no_grad(): + ty_output = tm.to(torch.float32)( + tx, position_ids=position_ids, position_embeddings=position_embeddings, + past_key_value=DynamicCache(), cache_position=cache_position, + ) + ty = ty_output[0] if isinstance(ty_output, tuple) else ty_output + + self._compare(jy, ty) + + def test_main_rms_norm(self): + shape = (self.batch_size, self.num_input_tokens, self.bonsai_config.hidden_size) + jx = jax.random.normal(jax.random.key(0), shape=shape) + tx = torch.tensor(np.array(jx, dtype=np.float32)) + jy = self.nnx_model.model.layers[0].input_layernorm(jx) + with torch.no_grad(): + ty = self.torch_model.model.layers[0].input_layernorm(tx) + self._compare(jy, ty) + + def test_main_self_attn(self): + shape = (self.batch_size, self.num_input_tokens, self.bonsai_config.hidden_size) + jx = jax.random.normal(jax.random.key(0), shape=shape) + tx = torch.tensor(np.array(jx, dtype=np.float32)) + + cache = self._init_cache(self.batch_size, self.num_input_tokens) + segment_ids = jnp.ones((self.batch_size, self.num_input_tokens), dtype=jnp.float32) + jy = self.nnx_model.model.layers[0].attn(jx, cache[0], segment_ids) + + cache_position = torch.arange(0, self.num_input_tokens, device=tx.device) + position_ids = cache_position.unsqueeze(0) + position_embeddings = self.torch_model.model.rotary_emb(tx, position_ids) + attention_mask = create_causal_mask( + config=self.torch_model.config, input_embeds=tx, attention_mask=None, + cache_position=cache_position, past_key_values=DynamicCache(), position_ids=position_ids, + ) + + with torch.no_grad(): + ty = self.torch_model.model.layers[0].self_attn.to(torch.float32)( + tx, position_ids=position_ids, position_embeddings=position_embeddings, + attention_mask=attention_mask, past_key_value=DynamicCache(), cache_position=cache_position, + )[0] + self._compare(jy, ty) + + def test_main_mlp(self): + shape = (self.batch_size, self.num_input_tokens, self.bonsai_config.hidden_size) + jx = jax.random.normal(jax.random.key(0), shape=shape) + tx = torch.tensor(np.array(jx, dtype=np.float32)) + jy = self.nnx_model.model.layers[0].mlp(jx) + with torch.no_grad(): + ty = self.torch_model.model.layers[0].mlp.to(torch.float32)(tx) + self._compare(jy, ty) + + def test_lm_head(self): + shape = (self.batch_size, self.num_input_tokens, self.bonsai_config.hidden_size) + jx = jax.random.normal(jax.random.key(0), shape=shape) + tx = torch.tensor(np.array(jx, dtype=np.float32)) + jy = self.nnx_model.lm_head(jx) + with torch.no_grad(): + ty = self.torch_model.lm_head.to(torch.float32)(tx) + self._compare(jy, ty) + + def test_local_transformer_lm_heads(self): + for ch in range(self.audio_channels): + shape = (self.batch_size, self.num_input_tokens, self.bonsai_config.local_dim) + jx = jax.random.normal(jax.random.key(ch), shape=shape) + tx = torch.tensor(np.array(jx, dtype=np.float32)) + jy = self.nnx_model.local_transformer_lm_heads[ch](jx) + with torch.no_grad(): + ty = self.torch_model.local_transformer_lm_heads[ch].to(torch.float32)(tx) + self._compare(jy, ty) + + def test_speech_group_downcast(self): + shape = (self.batch_size, self.num_input_tokens, self.bonsai_config.input_local_dim * self.group_size) + jx = jax.random.normal(jax.random.key(0), shape=shape) + tx = torch.tensor(np.array(jx, dtype=np.float32)) + jy = self.nnx_model.speech_group_downcast(jx) + with torch.no_grad(): + ty = self.torch_model.speech_group_downcast.to(torch.float32)(tx) + self._compare(jy, ty) + + def test_hidden_states_downcast(self): + shape = (self.batch_size, self.num_input_tokens, self.bonsai_config.hidden_size) + jx = jax.random.normal(jax.random.key(0), shape=shape) + tx = torch.tensor(np.array(jx, dtype=np.float32)) + jy = self.nnx_model.hidden_states_downcast(jx) + with torch.no_grad(): + ty = self.torch_model.hidden_states_downcast.to(torch.float32)(tx) + self._compare(jy, ty) + + def test_local_transformer_layer(self): + shape = (self.batch_size, self.num_input_tokens, self.bonsai_config.local_dim) + jx = jax.random.normal(jax.random.key(0), shape=shape) + tx = torch.tensor(np.array(jx, dtype=np.float32)) + + cache = self.nnx_model.local_transformer.init_cache( + cfg=self.nnx_model.local_qwen2_config, batch_size=self.batch_size, + token_len=self.num_input_tokens, generate_steps=0, dtype=jnp.float32 + ) + segment_ids = jnp.ones((self.batch_size, self.num_input_tokens)) + jy = self.nnx_model.local_transformer.layers[0](jx, cache[0], segment_ids) + + cache_position = torch.arange(0, self.num_input_tokens, device=tx.device) + position_ids = cache_position.unsqueeze(0) + position_embeddings = self.torch_model.local_transformer.rotary_emb(tx, position_ids) + + with torch.no_grad(): + ty_output = self.torch_model.local_transformer.layers[0].to(torch.float32)( + tx, position_ids=position_ids, position_embeddings=position_embeddings, + past_key_value=DynamicCache(), cache_position=cache_position, + ) + ty = ty_output[0] if isinstance(ty_output, tuple) else ty_output + + self._compare(jy, ty) + + def test_input_local_transformer_layer(self): + shape = (self.batch_size, self.num_input_tokens, self.bonsai_config.input_local_dim) + jx = jax.random.normal(jax.random.key(0), shape=shape) + tx = torch.tensor(np.array(jx, dtype=np.float32)) + + cache = self.nnx_model.input_local_transformer.init_cache( + cfg=self.nnx_model.input_local_qwen2_config, batch_size=self.batch_size, + token_len=self.num_input_tokens, generate_steps=0, dtype=jnp.float32 + ) + segment_ids = jnp.ones((self.batch_size, self.num_input_tokens)) + jy = self.nnx_model.input_local_transformer.layers[0](jx, cache[0], segment_ids) + + cache_position = torch.arange(0, self.num_input_tokens, device=tx.device) + position_ids = cache_position.unsqueeze(0) + position_embeddings = self.torch_model.input_local_transformer.rotary_emb(tx, position_ids) + attention_mask = torch.ones( + (self.batch_size, 1, self.num_input_tokens, self.num_input_tokens), + dtype=torch.float32, device=tx.device + ) + + with torch.no_grad(): + ty_output = self.torch_model.input_local_transformer.layers[0].to(torch.float32)( + tx, position_ids=position_ids, position_embeddings=position_embeddings, + attention_mask=attention_mask, past_key_value=DynamicCache(), cache_position=cache_position, + ) + ty = ty_output[0] if isinstance(ty_output, tuple) else ty_output + + self._compare(jy, ty) + + def test_apply_input_local_transformer(self): + shape = (self.batch_size, self.num_input_tokens // self.group_size, self.group_size, + self.bonsai_config.input_local_dim) + jx = jax.random.normal(jax.random.key(0), shape=shape) + tx = torch.tensor(np.array(jx, dtype=np.float32)) + jy = self.nnx_model.apply_input_local_transformer(jx, cache=None) + with torch.no_grad(): + ty = self.torch_model.apply_input_local_transformer(tx) + self._compare(jy, ty) + + def test_prepare_input_embeds(self): + num_groups = 3 + input_shape = (self.batch_size, self.audio_channels + 1, num_groups * self.group_size) + input_ids_np = np.random.randint(0, 1000, input_shape, dtype=np.int32) + for ch in range(self.audio_channels): + input_ids_np[:, ch + 1, :] = np.random.randint( + 0, self.nnx_model.speech_vocab_sizes[ch], (self.batch_size, num_groups * self.group_size) + ) + + input_ids_jax = jnp.array(input_ids_np) + input_ids_torch = torch.tensor(input_ids_np, dtype=torch.long) + + def text_embed_fn_jax(x): + return self.nnx_model.model.embedder.embedding.value[x] + jax_embeds = self.nnx_model._prepare_input_embeds(input_ids_jax, text_embed_fn_jax) + + with torch.no_grad(): + torch_embeds = self.torch_model._prepare_input_embeds(input_ids_torch) + + self._compare(jax_embeds, torch_embeds) + + def test_full_forward(self): + num_groups = 3 + input_shape = (self.batch_size, self.audio_channels + 1, num_groups * self.group_size) + input_ids_np = np.random.randint(0, 1000, input_shape, dtype=np.int32) + for ch in range(self.audio_channels): + input_ids_np[:, ch + 1, :] = np.random.randint( + 0, self.nnx_model.speech_vocab_sizes[ch], (self.batch_size, num_groups * self.group_size) + ) + + input_ids_jax = jnp.array(input_ids_np) + input_ids_torch = torch.tensor(input_ids_np, dtype=torch.long) + + cache_jax = self._init_cache(self.batch_size, num_groups) + text_logits_jax, local_hidden_jax, _ = self.nnx_model.forward(input_ids_jax, cache_jax) + + with torch.no_grad(): + attention_mask = torch.ones((self.batch_size, num_groups), dtype=torch.bool, device=input_ids_torch.device) + position_ids = torch.arange(num_groups).unsqueeze(0).expand(self.batch_size, -1) + cache_position = torch.arange(num_groups) + outputs_torch = self.torch_model( + input_ids=input_ids_torch, attention_mask=attention_mask, + position_ids=position_ids, cache_position=cache_position, + ) + text_logits_torch = outputs_torch.text_logits + local_hidden_torch = outputs_torch.local_hidden_states + + self._compare(text_logits_jax, text_logits_torch) + self._compare(local_hidden_jax, local_hidden_torch) + + +if __name__ == "__main__": + absltest.main() diff --git a/bonsai/models/qwen2/tests/run_model.py b/bonsai/models/qwen2/tests/run_model.py index 93519437..89c5030b 100644 --- a/bonsai/models/qwen2/tests/run_model.py +++ b/bonsai/models/qwen2/tests/run_model.py @@ -1,9 +1,9 @@ import jax -import os import jax.numpy as jnp import numpy as np from jax import P from jax._src.mesh import AxisType +from huggingface_hub import snapshot_download from transformers import AutoTokenizer from bonsai.models.qwen2 import modeling, params @@ -25,9 +25,9 @@ def tokenize(tokenizer, input: list[str], shd: P | None = None): def run_model(): - model_ckpt_path = os.path.expanduser("~/.cache/modelscope/hub/models/Qwen/Qwen2-7B") + model_name = "Qwen/Qwen2-7B" + model_ckpt_path = snapshot_download(model_name) - # Disable sharding - run on single GPU config = modeling.ModelConfig.qwen2_7b(use_sharding=False) # mesh, batch_shd = None, None @@ -40,8 +40,7 @@ def run_model(): ] tokenizer = AutoTokenizer.from_pretrained( - model_ckpt_path, - local_files_only=True, + model_name, trust_remote_code=True, ) From f65d52e6026695c39132a4176720260573fb1065 Mon Sep 17 00:00:00 2001 From: haibo Date: Tue, 20 Jan 2026 10:24:13 +0800 Subject: [PATCH 10/18] fix: audio encoder and audio decoder padding config --- bonsai/models/mimo_audio/mimo_audio_tokenizer.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/bonsai/models/mimo_audio/mimo_audio_tokenizer.py b/bonsai/models/mimo_audio/mimo_audio_tokenizer.py index 2b668b0d..bb4f89c4 100644 --- a/bonsai/models/mimo_audio/mimo_audio_tokenizer.py +++ b/bonsai/models/mimo_audio/mimo_audio_tokenizer.py @@ -504,7 +504,7 @@ def __init__(self, config: MiMoAudioTokenizerConfig, dtype=jnp.float32, rngs: Op in_features=config.n_mels, out_features=config.d_model, kernel_size=config.kernel_size, - padding="SAME", + padding=1, param_dtype=dtype, rngs=rngs ), @@ -516,7 +516,7 @@ def __init__(self, config: MiMoAudioTokenizerConfig, dtype=jnp.float32, rngs: Op out_features=config.d_model, kernel_size=config.kernel_size, strides=config.stride_size, - padding="SAME", + padding=1, param_dtype=dtype, rngs=rngs ), From fb4e25278d9b967f110d6784f522b68a50bdd1ee Mon Sep 17 00:00:00 2001 From: haibo Date: Tue, 20 Jan 2026 10:24:46 +0800 Subject: [PATCH 11/18] test: mimo audio --- bonsai/models/mimo_audio/test/test_outputs_mimo_audio.py | 3 +++ 1 file changed, 3 insertions(+) diff --git a/bonsai/models/mimo_audio/test/test_outputs_mimo_audio.py b/bonsai/models/mimo_audio/test/test_outputs_mimo_audio.py index 95001350..f8a09c28 100644 --- a/bonsai/models/mimo_audio/test/test_outputs_mimo_audio.py +++ b/bonsai/models/mimo_audio/test/test_outputs_mimo_audio.py @@ -15,6 +15,9 @@ from bonsai.models.mimo_audio import params from bonsai.models.mimo_audio.mimo_audio_configuration import MiMoAudioConfig, MiMoAudioArguments + +# Since the transformer library does not yet support mimo-audio-tokenizer, +# during testing, the official implementation code of mimo-audio-tokenizer (https://github.com/XiaomiMiMo/MiMo-Audio) needs to be copied to the corresponding location. from bonsai.models.mimo_audio.pytorch.src.mimo_audio.modeling_mimo_audio import ( MiMoAudioForCausalLM as TorchMiMoAudio, MiMoAudioConfig as TorchMiMoAudioConfig, From 710ce6dcc6592e53c28a7ada24bc6a4bcf851392 Mon Sep 17 00:00:00 2001 From: haibo Date: Tue, 20 Jan 2026 10:26:17 +0800 Subject: [PATCH 12/18] test: mimo_audio_tokenizer --- .../test/test_outputs_mimo_audio_tokenizer.py | 180 ++++++++++++++++++ 1 file changed, 180 insertions(+) create mode 100644 bonsai/models/mimo_audio/test/test_outputs_mimo_audio_tokenizer.py diff --git a/bonsai/models/mimo_audio/test/test_outputs_mimo_audio_tokenizer.py b/bonsai/models/mimo_audio/test/test_outputs_mimo_audio_tokenizer.py new file mode 100644 index 00000000..4afca147 --- /dev/null +++ b/bonsai/models/mimo_audio/test/test_outputs_mimo_audio_tokenizer.py @@ -0,0 +1,180 @@ +import os +import json +from absl.testing import absltest +import jax +import jax.numpy as jnp +import numpy as np +import torch +from flax import nnx +from huggingface_hub import snapshot_download + +from bonsai.models.mimo_audio.mimo_audio_tokenizer_params import load_tokenizer_weights_from_safetensors +from bonsai.models.mimo_audio.mimo_audio_tokenizer_configuration import MiMoAudioTokenizerConfig as JaxTokenizerConfig + +# Since the transformer library does not yet support mimo-audio-tokenizer, +# during testing, the official implementation code of mimo-audio-tokenizer (https://github.com/XiaomiMiMo/MiMo-Audio) needs to be copied to the corresponding location. +from bonsai.models.mimo_audio.pytorch.src.mimo_audio_tokenizer.modeling_audio_tokenizer import ( + MiMoAudioTokenizer as TorchTokenizer, +) +from bonsai.models.mimo_audio.pytorch.src.mimo_audio_tokenizer.configuration_audio_tokenizer import ( + MiMoAudioTokenizerConfig as TorchTokenizerConfig, +) + + +class TestMiMoAudioTokenizerOutputs(absltest.TestCase): + @classmethod + def setUpClass(cls): + super().setUpClass() + jax.config.update("jax_default_matmul_precision", "float32") + jax.config.update("jax_platforms", "cpu") + + model_name = "XiaomiMiMo/MiMo-Audio-Tokenizer" + model_ckpt_path = snapshot_download(model_name) + + safetensors_path = os.path.join(model_ckpt_path, "model.safetensors") + config_path = os.path.join(model_ckpt_path, "config.json") + + torch_config = TorchTokenizerConfig.from_pretrained(model_ckpt_path) + cls.torch_model = TorchTokenizer.from_pretrained( + model_ckpt_path, config=torch_config, torch_dtype=torch.float32 + ) + + cls.torch_model = cls.torch_model.cpu() + cls.torch_model.eval() + + def move_to_cpu(module): + for child in module.children(): + move_to_cpu(child) + for name, param in module._parameters.items(): + if param is not None: + module._parameters[name] = param.cpu() + for name, buf in module._buffers.items(): + if buf is not None: + module._buffers[name] = buf.cpu() + + move_to_cpu(cls.torch_model) + + with open(config_path) as f: + config_dict = json.load(f) + jax_config = JaxTokenizerConfig(**config_dict, use_sharding=False) + cls.nnx_model = load_tokenizer_weights_from_safetensors( + config=jax_config, safetensors_path=safetensors_path, + dtype=jnp.float32, mesh=None, rngs=nnx.Rngs(params=0) + ) + + cls.batch_size = 1 + cls.n_mels = jax_config.n_mels + cls.mel_frames = 50 + cls.d_model = jax_config.d_model + cls.tol = 1e-3 + + def _compare(self, jy, ty): + if ty.dim() == 2 and jy.ndim == 3: + ty = ty.unsqueeze(0) + torch.testing.assert_close( + torch.tensor(np.array(jy, dtype=np.float32)), ty, + rtol=self.tol, atol=self.tol, check_dtype=False, + ) + + def test_encoder_conv1(self): + mels = torch.randn(self.batch_size, self.n_mels, self.mel_frames, dtype=torch.float32) + jx = jnp.array(mels.permute(0, 2, 1).numpy()) + + jy = jax.nn.gelu(self.nnx_model.encoder.conv1(jx)) + ty = torch.nn.functional.gelu(self.torch_model.encoder.conv1(mels)) + + self._compare(jy, ty.permute(0, 2, 1)) + + def test_encoder_conv2(self): + x = torch.randn(self.batch_size, self.d_model, self.mel_frames, dtype=torch.float32) + jx = jnp.array(x.permute(0, 2, 1).numpy()) + + jy = jax.nn.gelu(self.nnx_model.encoder.conv2(jx)) + ty = torch.nn.functional.gelu(self.torch_model.encoder.conv2(x)) + + self._compare(jy, ty.permute(0, 2, 1)) + + def test_quantizer_encode(self): + seq_len = 12 + x = torch.randn(seq_len, self.d_model, dtype=torch.float32) + jx = jnp.array(x.numpy()) + + jcodes, jquantized = self.nnx_model.encoder.quantizer.encode(jx, mask=None, n_q=None) + + tcodes = self.torch_model.encoder.quantizer.encode(x.unsqueeze(0)) + + np.testing.assert_array_equal(np.array(jcodes), tcodes.squeeze(1).numpy()) + + def test_quantizer_decode(self): + num_q = self.nnx_model.encoder.quantizer.n_q + seq_len = 12 + + codes_list = [] + for i in range(num_q): + codebook_size = self.nnx_model.encoder.quantizer.codebooks[i].value.shape[0] + codes_list.append(torch.randint(0, codebook_size, (seq_len,), dtype=torch.long)) + codes = torch.stack(codes_list, dim=0) + jcodes = jnp.array(codes.numpy()) + + jdecoded = self.nnx_model.encoder.decode_vq(jcodes) + tdecoded = self.torch_model.encoder.decode_vq(codes) + + self._compare(jdecoded, tdecoded) + + def test_decoder_dconv1(self): + seq_len = 12 + x = torch.randn(self.batch_size, seq_len, self.d_model, dtype=torch.float32) + jx = jnp.array(x.numpy()) + + input_length = torch.tensor([seq_len], dtype=torch.long) + jinput_length = jnp.array([seq_len], dtype=jnp.int32) + + if self.nnx_model.decoder.dconv1 is not None: + jy, jout_len = self.nnx_model.decoder.dconv1(jx, jinput_length) + ty, tout_len = self.torch_model.decoder.dconv1(x, input_length, output_dim=3) + + self._compare(jy, ty) + np.testing.assert_array_equal(np.array(jout_len), tout_len.numpy()) + + def test_decoder_dconv2(self): + seq_len = 24 + x = torch.randn(self.batch_size, seq_len, self.d_model, dtype=torch.float32) + jx = jnp.array(x.numpy()) + + input_length = torch.tensor([seq_len], dtype=torch.long) + jinput_length = jnp.array([seq_len], dtype=jnp.int32) + + jy, jout_len = self.nnx_model.decoder.dconv2(jx, jinput_length) + + tx = torch.masked_select(x, torch.ones_like(x, dtype=bool)).view(-1, self.d_model) + ty, tout_len = self.torch_model.decoder.dconv2(tx, input_length, output_dim=3) + + self._compare(jy, ty) + np.testing.assert_array_equal(np.array(jout_len), tout_len.numpy()) + + def test_vocoder_embeddings(self): + seq_len = 48 + x = torch.randn(self.batch_size, seq_len, self.n_mels, dtype=torch.float32) + jx = jnp.array(x.numpy()) + + jy = self.nnx_model.decoder.vocoder.embeddings(jx) + ty = self.torch_model.decoder.vocoder.embeddings(x) + + self._compare(jy, ty) + + def test_vocoder_istft_head(self): + vocoder_dim = self.nnx_model.decoder.vocoder.config.vocoder_dim + seq_len = 48 + x = torch.randn(self.batch_size, seq_len, vocoder_dim, dtype=torch.float32) + jx = jnp.array(x.numpy()) + + jy = self.nnx_model.decoder.vocoder.head(jx) + ty = self.torch_model.decoder.vocoder.head(x) + + torch.testing.assert_close( + torch.tensor(np.array(jy, dtype=np.float32)), ty, + rtol=1e-2, atol=1e-2, check_dtype=False, + ) + +if __name__ == "__main__": + absltest.main() From 27c7553e6feb21cbf280913315ce81c5e9bb6f6b Mon Sep 17 00:00:00 2001 From: haibo Date: Tue, 20 Jan 2026 11:25:11 +0800 Subject: [PATCH 13/18] pre-commit fixes --- bonsai/models/mimo_audio/__init__.py | 4 +- .../mimo_audio/mimo_audio_configuration.py | 16 +- .../models/mimo_audio/mimo_audio_tokenizer.py | 362 +++++++++++------- .../mimo_audio_tokenizer_configuration.py | 113 +++--- .../mimo_audio/mimo_audio_tokenizer_params.py | 267 +++++++++---- bonsai/models/mimo_audio/modeling.py | 248 +++--------- bonsai/models/mimo_audio/params.py | 64 +++- bonsai/models/mimo_audio/test/run_model.py | 35 +- .../test/test_outputs_mimo_audio.py | 117 ++++-- .../test/test_outputs_mimo_audio_tokenizer.py | 18 +- bonsai/models/qwen2/modeling.py | 47 +-- bonsai/models/qwen2/params.py | 9 +- bonsai/models/qwen2/tests/__init__.py | 2 +- bonsai/models/qwen2/tests/run_model.py | 24 +- .../models/qwen2/tests/test_outputs_qwen2.py | 7 +- 15 files changed, 704 insertions(+), 629 deletions(-) diff --git a/bonsai/models/mimo_audio/__init__.py b/bonsai/models/mimo_audio/__init__.py index aefeeb45..5f102705 100644 --- a/bonsai/models/mimo_audio/__init__.py +++ b/bonsai/models/mimo_audio/__init__.py @@ -1,4 +1,3 @@ -from bonsai.models.mimo_audio.mimo_audio import MimoAudio from bonsai.models.mimo_audio.modeling import ( MiMoAudioConfig, MiMoAudioArguments, @@ -11,11 +10,10 @@ ) __all__ = [ - "MimoAudio", "MiMoAudioConfig", "MiMoAudioArguments", "FlaxMiMoAudioForCausalLM", "FlaxMiMoAudioTokenizer", "MiMoAudioTokenizerConfig", "MelSpectrogram", -] \ No newline at end of file +] diff --git a/bonsai/models/mimo_audio/mimo_audio_configuration.py b/bonsai/models/mimo_audio/mimo_audio_configuration.py index 87fe4389..75ec1811 100644 --- a/bonsai/models/mimo_audio/mimo_audio_configuration.py +++ b/bonsai/models/mimo_audio/mimo_audio_configuration.py @@ -1,5 +1,5 @@ from dataclasses import dataclass -from typing import Optional, TYPE_CHECKING +from typing import TYPE_CHECKING from bonsai.models.qwen3.modeling import ShardingCfg if TYPE_CHECKING: @@ -8,8 +8,6 @@ @dataclass class MiMoAudioConfig: - - # 主 Transformer 配置 vocab_size: int = 151680 hidden_size: int = 4096 num_hidden_layers: int = 36 @@ -23,24 +21,21 @@ class MiMoAudioConfig: group_size: int = 4 audio_channels: int = 8 - # Local Transformer config local_dim: int = 1024 local_layers: int = 16 local_attn_heads: int = 64 local_ffn_dim: int = 4096 local_attn_dropout: float = 0.1 - # Input Local Transformer config input_local_layers: int = 6 input_local_dim: int = 1024 input_full_attention: bool = True - # Sharding config shd_cfg: ShardingCfg = ShardingCfg.no_sharding() @classmethod def with_sharding(cls, **kwargs): - kwargs['shd_cfg'] = ShardingCfg.default() + kwargs["shd_cfg"] = ShardingCfg.default() return cls(**kwargs) def create_qwen2_config(self) -> "Qwen2Config": @@ -80,14 +75,13 @@ def create_local_qwen2_config(self) -> "Qwen2Config": def create_input_local_qwen2_config(self) -> "Qwen2Config": from bonsai.models.qwen2.modeling import ModelConfig as Qwen2Config - # input_full_attention=True -> use_causal_mask=False (bidirectional attention) return Qwen2Config( num_layers=6, vocab_size=151680, emb_dim=1024, - mlp_dim=4096, # 1024 * 4 + mlp_dim=4096, num_heads=64, - head_dim=16, # 1024 // 64 = 16 + head_dim=16, num_kv_heads=64, rope_theta=640000, norm_eps=1e-6, @@ -99,7 +93,6 @@ def create_input_local_qwen2_config(self) -> "Qwen2Config": @dataclass class MiMoAudioArguments: - """Arguments for special token indices""" model_name_or_path: str sosp_idx: int eosp_idx: int @@ -111,7 +104,6 @@ class MiMoAudioArguments: @dataclass class MiMoSamplerConfig: - """Sampler configuration for text/audio generation""" do_sample: bool = True temperature: float = 1.0 top_k: int = 50 diff --git a/bonsai/models/mimo_audio/mimo_audio_tokenizer.py b/bonsai/models/mimo_audio/mimo_audio_tokenizer.py index bb4f89c4..6bc88f2a 100644 --- a/bonsai/models/mimo_audio/mimo_audio_tokenizer.py +++ b/bonsai/models/mimo_audio/mimo_audio_tokenizer.py @@ -47,19 +47,17 @@ def apply_rotary(x: Array, cos: Array, sin: Array) -> Array: class MelSpectrogram: - """Mel spectrogram computation for audio processing.""" - def __init__( - self, - sample_rate: int, - n_fft: int, - hop_length: int, - win_length: int, - f_min: float, - f_max: float, - n_mels: int, - power: float = 1.0, - center: bool = True, + self, + sample_rate: int, + n_fft: int, + hop_length: int, + win_length: int, + f_min: float, + f_max: float, + n_mels: int, + power: float = 1.0, + center: bool = True, ) -> None: self.sample_rate = int(sample_rate) self.n_fft = int(n_fft) @@ -132,9 +130,7 @@ def _frame_signal(self, waveform: jnp.ndarray) -> jnp.ndarray: num_frames = 1 starts = [idx * self.hop_length for idx in range(num_frames)] - frames = jnp.stack( - [waveform[start: start + frame_length] for start in starts], axis=0 - ) + frames = jnp.stack([waveform[start : start + frame_length] for start in starts], axis=0) return frames def _mel_spectrogram(self, waveform: jnp.ndarray) -> jnp.ndarray: @@ -200,11 +196,16 @@ def __call__(self, hidden_states: Array, position_ids: Array) -> Tuple[Array, Ar class ConvTranspose1d(nnx.Module): - """Custom 1D transposed convolution for specific audio processing requirements.""" - - def __init__(self, in_channels: int, out_channels: int, kernel_size: int, stride: int, - shd_cfg: MiMoShardingCfg | None = None, - dtype=jnp.float32, rngs: Optional[nnx.Rngs] = None): + def __init__( + self, + in_channels: int, + out_channels: int, + kernel_size: int, + stride: int, + shd_cfg: MiMoShardingCfg | None = None, + dtype=jnp.float32, + rngs: Optional[nnx.Rngs] = None, + ): self.stride = stride self.shd_cfg = shd_cfg or MiMoShardingCfg.no_sharding() @@ -230,8 +231,8 @@ def __call__(self, x: Array) -> Array: lhs=lhs, rhs=rhs, window_strides=(1,), - padding='VALID', - dimension_numbers=('NCH', 'OIH', 'NCH'), + padding="VALID", + dimension_numbers=("NCH", "OIH", "NCH"), ) y = y + self.bias.value[None, :, None] y = jnp.swapaxes(y, 1, 2) @@ -239,18 +240,22 @@ def __call__(self, x: Array) -> Array: class ISTFT(nnx.Module): - def __init__(self, n_fft: int, hop_length: int, win_length: int, padding: str = "same", - shd_cfg: MiMoShardingCfg | None = None, dtype=jnp.float32): + def __init__( + self, + n_fft: int, + hop_length: int, + win_length: int, + padding: str = "same", + shd_cfg: MiMoShardingCfg | None = None, + dtype=jnp.float32, + ): self.n_fft = n_fft self.hop_length = hop_length self.win_length = win_length self.padding = padding self.shd_cfg = shd_cfg or MiMoShardingCfg.no_sharding() - self.window = shard( - nnx.Param(jnp.hanning(win_length).astype(dtype)), - self.shd_cfg.istft_window - ) + self.window = shard(nnx.Param(jnp.hanning(win_length).astype(dtype)), self.shd_cfg.istft_window) self.pad = (self.win_length - self.hop_length) // 2 if padding == "same" else 0 @@ -286,26 +291,31 @@ def body(i, carry): audio, env = jax.lax.fori_loop(0, num_frames, body, (audio, env)) if self.pad > 0: - audio = audio[:, self.pad: -self.pad] - env = env[:, self.pad: -self.pad] + audio = audio[:, self.pad : -self.pad] + env = env[:, self.pad : -self.pad] env = jnp.maximum(env, 1e-11) audio = audio / env return audio class ISTFTHead(nnx.Module): - def __init__(self, dim: int, n_fft: int, hop_length: int, padding: str = "same", - shd_cfg: MiMoShardingCfg | None = None, - dtype=jnp.float32, rngs: Optional[nnx.Rngs] = None): + def __init__( + self, + dim: int, + n_fft: int, + hop_length: int, + padding: str = "same", + shd_cfg: MiMoShardingCfg | None = None, + dtype=jnp.float32, + rngs: Optional[nnx.Rngs] = None, + ): self.shd_cfg = shd_cfg or MiMoShardingCfg.no_sharding() - self.linear = shard( - nnx.Linear(dim, n_fft + 2, dtype=dtype, rngs=rngs), - self.shd_cfg.istft_linear_weight - ) + self.linear = shard(nnx.Linear(dim, n_fft + 2, dtype=dtype, rngs=rngs), self.shd_cfg.istft_linear_weight) - self.istft = ISTFT(n_fft=n_fft, hop_length=hop_length, win_length=n_fft, - padding=padding, shd_cfg=self.shd_cfg, dtype=dtype) + self.istft = ISTFT( + n_fft=n_fft, hop_length=hop_length, win_length=n_fft, padding=padding, shd_cfg=self.shd_cfg, dtype=dtype + ) def __call__(self, hidden_states: Array) -> Array: x = self.linear(hidden_states) @@ -327,9 +337,16 @@ def __call__(self, hidden_states: Array) -> Array: class Attention(nnx.Module): - def __init__(self, embed_dim: int, num_heads: int, window_size: Tuple[int, int], causal: bool, - shd_cfg: MiMoShardingCfg, dtype=jnp.float32, - rngs: Optional[nnx.Rngs] = None): + def __init__( + self, + embed_dim: int, + num_heads: int, + window_size: Tuple[int, int], + causal: bool, + shd_cfg: MiMoShardingCfg, + dtype=jnp.float32, + rngs: Optional[nnx.Rngs] = None, + ): self.embed_dim = embed_dim self.num_heads = num_heads self.head_dim = embed_dim // num_heads @@ -339,21 +356,15 @@ def __init__(self, embed_dim: int, num_heads: int, window_size: Tuple[int, int], self.shd_cfg = shd_cfg self.q_proj = shard( - nnx.Linear(embed_dim, embed_dim, use_bias=True, dtype=dtype, rngs=rngs), - shd_cfg.attn_qkvo_weight + nnx.Linear(embed_dim, embed_dim, use_bias=True, dtype=dtype, rngs=rngs), shd_cfg.attn_qkvo_weight ) self.k_proj = shard( - nnx.Linear(embed_dim, embed_dim, use_bias=False, dtype=dtype, rngs=rngs), - shd_cfg.attn_qkvo_weight + nnx.Linear(embed_dim, embed_dim, use_bias=False, dtype=dtype, rngs=rngs), shd_cfg.attn_qkvo_weight ) self.v_proj = shard( - nnx.Linear(embed_dim, embed_dim, use_bias=True, dtype=dtype, rngs=rngs), - shd_cfg.attn_qkvo_weight - ) - self.out_proj = shard( - nnx.Linear(embed_dim, embed_dim, dtype=dtype, rngs=rngs), - shd_cfg.attn_qkvo_weight + nnx.Linear(embed_dim, embed_dim, use_bias=True, dtype=dtype, rngs=rngs), shd_cfg.attn_qkvo_weight ) + self.out_proj = shard(nnx.Linear(embed_dim, embed_dim, dtype=dtype, rngs=rngs), shd_cfg.attn_qkvo_weight) def _window_mask(self, seq_len: int) -> Optional[Array]: left, right = self.window_size @@ -409,32 +420,31 @@ def reshape(t): class TransformerLayer(nnx.Module): - def __init__(self, d_model: int, attention_heads: int, ffn_dim: int, causal: bool, - attn_window_size: Tuple[int, int], shd_cfg: MiMoShardingCfg, dtype=jnp.float32, - rngs: Optional[nnx.Rngs] = None): + def __init__( + self, + d_model: int, + attention_heads: int, + ffn_dim: int, + causal: bool, + attn_window_size: Tuple[int, int], + shd_cfg: MiMoShardingCfg, + dtype=jnp.float32, + rngs: Optional[nnx.Rngs] = None, + ): self.act = jax.nn.gelu self.shd_cfg = shd_cfg - self.self_attn = Attention(d_model, attention_heads, attn_window_size, causal, - shd_cfg, dtype=dtype, rngs=rngs) + self.self_attn = Attention(d_model, attention_heads, attn_window_size, causal, shd_cfg, dtype=dtype, rngs=rngs) self.self_attn_layer_norm = shard( - nnx.LayerNorm(d_model, epsilon=1e-6, param_dtype=dtype, rngs=rngs), - shd_cfg.norm_scale + nnx.LayerNorm(d_model, epsilon=1e-6, param_dtype=dtype, rngs=rngs), shd_cfg.norm_scale ) self.final_layer_norm = shard( - nnx.LayerNorm(d_model, epsilon=1e-6, param_dtype=dtype, rngs=rngs), - shd_cfg.norm_scale + nnx.LayerNorm(d_model, epsilon=1e-6, param_dtype=dtype, rngs=rngs), shd_cfg.norm_scale ) - self.fc1 = shard( - nnx.Linear(d_model, ffn_dim, dtype=dtype, rngs=rngs), - shd_cfg.ffn_weight_in - ) - self.fc2 = shard( - nnx.Linear(ffn_dim, d_model, dtype=dtype, rngs=rngs), - shd_cfg.ffn_weight_out - ) + self.fc1 = shard(nnx.Linear(d_model, ffn_dim, dtype=dtype, rngs=rngs), shd_cfg.ffn_weight_in) + self.fc2 = shard(nnx.Linear(ffn_dim, d_model, dtype=dtype, rngs=rngs), shd_cfg.ffn_weight_out) def __call__(self, hidden_states: Array, mask: Optional[Array], rope: Optional[Tuple[Array, Array]]) -> Array: residual = hidden_states @@ -451,9 +461,15 @@ def __call__(self, hidden_states: Array, mask: Optional[Array], rope: Optional[T class ResidualVectorQuantizer(nnx.Module): - def __init__(self, dimension: int, n_q: int, bins: Sequence[int], - shd_cfg: MiMoShardingCfg, dtype=jnp.float32, - rngs: Optional[nnx.Rngs] = None): + def __init__( + self, + dimension: int, + n_q: int, + bins: Sequence[int], + shd_cfg: MiMoShardingCfg, + dtype=jnp.float32, + rngs: Optional[nnx.Rngs] = None, + ): self.dimension = dimension self.n_q = n_q self.shd_cfg = shd_cfg @@ -465,8 +481,9 @@ def __init__(self, dimension: int, n_q: int, bins: Sequence[int], codebooks_list.append(shard(nnx.Param(embed), shd_cfg.codebook)) self.codebooks = nnx.List(codebooks_list) - def encode(self, hidden_states: Array, mask: Optional[Array] = None, n_q: Optional[int] = None) -> Tuple[ - Array, Array]: + def encode( + self, hidden_states: Array, mask: Optional[Array] = None, n_q: Optional[int] = None + ) -> Tuple[Array, Array]: num_levels = n_q or self.n_q residual = hidden_states quantized = jnp.zeros_like(hidden_states) @@ -506,9 +523,9 @@ def __init__(self, config: MiMoAudioTokenizerConfig, dtype=jnp.float32, rngs: Op kernel_size=config.kernel_size, padding=1, param_dtype=dtype, - rngs=rngs + rngs=rngs, ), - self.shd_cfg.conv_weight + self.shd_cfg.conv_weight, ) self.conv2 = shard( nnx.Conv( @@ -518,25 +535,37 @@ def __init__(self, config: MiMoAudioTokenizerConfig, dtype=jnp.float32, rngs: Op strides=config.stride_size, padding=1, param_dtype=dtype, - rngs=rngs + rngs=rngs, ), - self.shd_cfg.conv_weight + self.shd_cfg.conv_weight, ) - self.position_embedding = RotaryEmbedding(config.rope_theta, config.d_model // config.encoder_attention_heads, - config.max_audio_seconds * config.sampling_rate // config.hop_length, - config.rope_type, dtype=dtype) + self.position_embedding = RotaryEmbedding( + config.rope_theta, + config.d_model // config.encoder_attention_heads, + config.max_audio_seconds * config.sampling_rate // config.hop_length, + config.rope_type, + dtype=dtype, + ) - self.layers = nnx.List([ - TransformerLayer(config.d_model, config.encoder_attention_heads, config.encoder_ffn_dim, - config.encoder_causal, tuple(config.encoder_attn_window_size), - self.shd_cfg, dtype=dtype, rngs=rngs) - for _ in range(config.encoder_layers) - ]) + self.layers = nnx.List( + [ + TransformerLayer( + config.d_model, + config.encoder_attention_heads, + config.encoder_ffn_dim, + config.encoder_causal, + tuple(config.encoder_attn_window_size), + self.shd_cfg, + dtype=dtype, + rngs=rngs, + ) + for _ in range(config.encoder_layers) + ] + ) self.layer_norm = shard( - nnx.LayerNorm(config.d_model, epsilon=1e-6, param_dtype=dtype, rngs=rngs), - self.shd_cfg.norm_scale + nnx.LayerNorm(config.d_model, epsilon=1e-6, param_dtype=dtype, rngs=rngs), self.shd_cfg.norm_scale ) if config.avg_pooler != 1: @@ -549,13 +578,12 @@ def __init__(self, config: MiMoAudioTokenizerConfig, dtype=jnp.float32, rngs: Op padding="SAME", use_bias=False, param_dtype=dtype, - rngs=rngs + rngs=rngs, ), - self.shd_cfg.conv_weight + self.shd_cfg.conv_weight, ) self.down_norm = shard( - nnx.LayerNorm(config.d_model, epsilon=1e-6, param_dtype=dtype, rngs=rngs), - self.shd_cfg.norm_scale + nnx.LayerNorm(config.d_model, epsilon=1e-6, param_dtype=dtype, rngs=rngs), self.shd_cfg.norm_scale ) else: self.down_sample_layer = None @@ -563,8 +591,9 @@ def __init__(self, config: MiMoAudioTokenizerConfig, dtype=jnp.float32, rngs: Op if config.num_quantizers: bins = config.codebook_size or [1024] - self.quantizer = ResidualVectorQuantizer(config.d_model, config.num_quantizers, bins, - self.shd_cfg, dtype=dtype, rngs=rngs) + self.quantizer = ResidualVectorQuantizer( + config.d_model, config.num_quantizers, bins, self.shd_cfg, dtype=dtype, rngs=rngs + ) else: self.quantizer = None @@ -572,8 +601,9 @@ def get_output_length(self, mel_len: Array) -> Array: tgt = mel_len + 3 - self.config.kernel_size return (tgt + 2 - self.config.kernel_size) // self.config.stride_size + 1 - def __call__(self, input_features: Array, input_lens: Array, use_quantizer: bool = True, - n_q: Optional[int] = None) -> EncoderOutput: + def __call__( + self, input_features: Array, input_lens: Array, use_quantizer: bool = True, n_q: Optional[int] = None + ) -> EncoderOutput: x = input_features x = jax.nn.gelu(self.conv1(x)) x = shard(x, self.shd_cfg.act_btd) @@ -598,7 +628,8 @@ def __call__(self, input_features: Array, input_lens: Array, use_quantizer: bool x = shard(x, self.shd_cfg.act_btd) lengths = (lengths // self.config.avg_pooler) + ((lengths % self.config.avg_pooler) != 0).astype( - lengths.dtype) + lengths.dtype + ) max_len = x.shape[1] mask = make_sequence_mask(lengths, max_len) x = self.down_norm(x) @@ -619,18 +650,25 @@ def decode_vq(self, codes: Array) -> Array: class CausalConvTranspose1d(nnx.Module): - def __init__(self, in_channels: int, out_channels: int, kernel_size: int, stride: int, - shd_cfg: MiMoShardingCfg | None = None, - dtype=jnp.float32, rngs: Optional[nnx.Rngs] = None): + def __init__( + self, + in_channels: int, + out_channels: int, + kernel_size: int, + stride: int, + shd_cfg: MiMoShardingCfg | None = None, + dtype=jnp.float32, + rngs: Optional[nnx.Rngs] = None, + ): self.shd_cfg = shd_cfg or MiMoShardingCfg.no_sharding() - self.conv = ConvTranspose1d(in_channels, out_channels, kernel_size, stride, - shd_cfg=self.shd_cfg, dtype=dtype, rngs=rngs) + self.conv = ConvTranspose1d( + in_channels, out_channels, kernel_size, stride, shd_cfg=self.shd_cfg, dtype=dtype, rngs=rngs + ) self.norm = shard( - nnx.GroupNorm(num_features=out_channels, num_groups=1, epsilon=1e-5, - param_dtype=dtype, rngs=rngs), - self.shd_cfg.norm_scale + nnx.GroupNorm(num_features=out_channels, num_groups=1, epsilon=1e-5, param_dtype=dtype, rngs=rngs), + self.shd_cfg.norm_scale, ) self.kernel_size = kernel_size @@ -653,29 +691,46 @@ def __init__(self, config: MiMoAudioTokenizerConfig, dtype=jnp.float32, rngs: Op self.embeddings = shard( nnx.Linear(config.n_mels, config.vocoder_dim, use_bias=False, dtype=dtype, rngs=rngs), - self.shd_cfg.attn_qkvo_weight + self.shd_cfg.attn_qkvo_weight, ) - self.position_embedding = RotaryEmbedding(config.rope_theta, - config.vocoder_dim // config.vocoder_attention_heads, - config.max_audio_seconds * config.sampling_rate // config.hop_length, - config.rope_type, dtype=dtype) + self.position_embedding = RotaryEmbedding( + config.rope_theta, + config.vocoder_dim // config.vocoder_attention_heads, + config.max_audio_seconds * config.sampling_rate // config.hop_length, + config.rope_type, + dtype=dtype, + ) - self.layers = nnx.List([ - TransformerLayer(config.vocoder_dim, config.vocoder_attention_heads, config.vocoder_intermediate_dim, - False, tuple(config.vocoder_attn_window_size), - self.shd_cfg, dtype=dtype, rngs=rngs) - for _ in range(config.vocoder_num_layers) - ]) + self.layers = nnx.List( + [ + TransformerLayer( + config.vocoder_dim, + config.vocoder_attention_heads, + config.vocoder_intermediate_dim, + False, + tuple(config.vocoder_attn_window_size), + self.shd_cfg, + dtype=dtype, + rngs=rngs, + ) + for _ in range(config.vocoder_num_layers) + ] + ) self.layer_norm = shard( - nnx.LayerNorm(config.vocoder_dim, epsilon=1e-6, param_dtype=dtype, rngs=rngs), - self.shd_cfg.norm_scale + nnx.LayerNorm(config.vocoder_dim, epsilon=1e-6, param_dtype=dtype, rngs=rngs), self.shd_cfg.norm_scale ) - self.head = ISTFTHead(config.vocoder_dim, config.nfft, config.hop_length, - config.vocoder_padding, shd_cfg=self.shd_cfg, - dtype=dtype, rngs=rngs) + self.head = ISTFTHead( + config.vocoder_dim, + config.nfft, + config.hop_length, + config.vocoder_padding, + shd_cfg=self.shd_cfg, + dtype=dtype, + rngs=rngs, + ) def __call__(self, mels: Array, input_length: Array) -> VocoderOutput: x = self.embeddings(mels) @@ -698,31 +753,55 @@ def __init__(self, config: MiMoAudioTokenizerConfig, dtype=jnp.float32, rngs: Op self.shd_cfg = config.shd_cfg if config.avg_pooler != 1: - self.dconv1 = CausalConvTranspose1d(config.d_model, config.d_model, config.avg_pooler, - config.avg_pooler, shd_cfg=self.shd_cfg, - dtype=dtype, rngs=rngs) + self.dconv1 = CausalConvTranspose1d( + config.d_model, + config.d_model, + config.avg_pooler, + config.avg_pooler, + shd_cfg=self.shd_cfg, + dtype=dtype, + rngs=rngs, + ) else: self.dconv1 = None - self.position_embedding = RotaryEmbedding(config.rope_theta, config.d_model // config.decoder_attention_heads, - config.max_audio_seconds * config.sampling_rate // config.hop_length, - config.rope_type, dtype=dtype) + self.position_embedding = RotaryEmbedding( + config.rope_theta, + config.d_model // config.decoder_attention_heads, + config.max_audio_seconds * config.sampling_rate // config.hop_length, + config.rope_type, + dtype=dtype, + ) - self.layers = nnx.List([ - TransformerLayer(config.d_model, config.decoder_attention_heads, config.decoder_ffn_dim, - config.decoder_causal, tuple(config.decoder_attn_window_size), - self.shd_cfg, dtype=dtype, rngs=rngs) - for _ in range(config.decoder_layers) - ]) + self.layers = nnx.List( + [ + TransformerLayer( + config.d_model, + config.decoder_attention_heads, + config.decoder_ffn_dim, + config.decoder_causal, + tuple(config.decoder_attn_window_size), + self.shd_cfg, + dtype=dtype, + rngs=rngs, + ) + for _ in range(config.decoder_layers) + ] + ) self.layer_norm = shard( - nnx.LayerNorm(config.d_model, epsilon=1e-6, param_dtype=dtype, rngs=rngs), - self.shd_cfg.norm_scale + nnx.LayerNorm(config.d_model, epsilon=1e-6, param_dtype=dtype, rngs=rngs), self.shd_cfg.norm_scale ) - self.dconv2 = CausalConvTranspose1d(config.d_model, config.n_mels, config.decoder_kernel_size, - config.decoder_stride_size, shd_cfg=self.shd_cfg, - dtype=dtype, rngs=rngs) + self.dconv2 = CausalConvTranspose1d( + config.d_model, + config.n_mels, + config.decoder_kernel_size, + config.decoder_stride_size, + shd_cfg=self.shd_cfg, + dtype=dtype, + rngs=rngs, + ) self.vocoder = TransformerVocos(config, dtype=dtype, rngs=rngs) @@ -753,8 +832,9 @@ def __call__(self, mels: Array, input_lens: Array, use_quantizer: bool = True) - enc = self.encoder(mels, input_lens, use_quantizer=use_quantizer) return self.decoder(enc.hidden_states, enc.output_lengths) - def encode(self, mels: Array, input_lens: Array, use_quantizer: bool = True, - n_q: Optional[int] = None) -> EncoderOutput: + def encode( + self, mels: Array, input_lens: Array, use_quantizer: bool = True, n_q: Optional[int] = None + ) -> EncoderOutput: return self.encoder(mels, input_lens, use_quantizer=use_quantizer, n_q=n_q) def decode(self, codes: Array) -> Array: diff --git a/bonsai/models/mimo_audio/mimo_audio_tokenizer_configuration.py b/bonsai/models/mimo_audio/mimo_audio_tokenizer_configuration.py index 2512fc27..111d9b9b 100644 --- a/bonsai/models/mimo_audio/mimo_audio_tokenizer_configuration.py +++ b/bonsai/models/mimo_audio/mimo_audio_tokenizer_configuration.py @@ -12,11 +12,6 @@ @dataclass(slots=True, frozen=True) class MiMoShardingCfg: - """Sharding configuration for MiMo Audio Tokenizer. - - Controls how model parameters and activations are distributed across devices. - """ - # Conv layer weight sharding conv_weight: ShardingSpec # (in_channels, out_channels, kernel_size) conv_bias: ShardingSpec # (out_channels,) @@ -106,51 +101,51 @@ class MiMoAudioTokenizerConfig(PretrainedConfig): model_type = "mimo_audio_tokenizer" def __init__( - self, - max_audio_seconds: int = 1800, - stride_size: int = 2, - avg_pooler: int = 2, - d_model: int = 1280, - scale_embedding: bool = False, - kernel_size: int = 3, - activation_function: str = "gelu", - encoder_layers: int = 32, - encoder_skip_layer_id: int = 3, - encoder_attention_heads: int = 20, - encoder_ffn_dim: int = 5120, - encoder_causal: bool = False, - encoder_attn_window_size: list[int] = None, # [-1,-1] - decoder_layers: int = 32, - decoder_attention_heads: int = 20, - decoder_ffn_dim: int = 5120, - decoder_kernel_size: int = 3, - decoder_stride_size: int = 2, - decoder_causal: bool = True, - decoder_attn_window_size: list[int] = None, # [-1,-1] - nfft: int = 960, - vocoder_dim: int = 256, - vocoder_intermediate_dim: int = 1024, - vocoder_num_layers: int = 16, - n_mels: int = 128, - sampling_rate: int = 24000, - hop_length: int = 240, - window_size: int = 960, - vocoder_padding: str = "same", - fmin: int = 0, - fmax: int = None, - num_quantizers: int = 20, - codebook_size: list[int] = None, - # [1024,1024,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128] - threshold_ema_dead_code: int = 2, - position_embedding_type: str = "rope", - rope_theta: int = 10000, - rope_type: str = "default", - ln_type: str = "LayerNorm", - vocoder_attention_heads: int = 16, - vocoder_attn_window_size: list[int] = None, # [40,10] - use_sharding: bool = False, - shd_cfg: MiMoShardingCfg | None = None, - **kwargs, + self, + max_audio_seconds: int = 1800, + stride_size: int = 2, + avg_pooler: int = 2, + d_model: int = 1280, + scale_embedding: bool = False, + kernel_size: int = 3, + activation_function: str = "gelu", + encoder_layers: int = 32, + encoder_skip_layer_id: int = 3, + encoder_attention_heads: int = 20, + encoder_ffn_dim: int = 5120, + encoder_causal: bool = False, + encoder_attn_window_size: list[int] = None, # [-1,-1] + decoder_layers: int = 32, + decoder_attention_heads: int = 20, + decoder_ffn_dim: int = 5120, + decoder_kernel_size: int = 3, + decoder_stride_size: int = 2, + decoder_causal: bool = True, + decoder_attn_window_size: list[int] = None, # [-1,-1] + nfft: int = 960, + vocoder_dim: int = 256, + vocoder_intermediate_dim: int = 1024, + vocoder_num_layers: int = 16, + n_mels: int = 128, + sampling_rate: int = 24000, + hop_length: int = 240, + window_size: int = 960, + vocoder_padding: str = "same", + fmin: int = 0, + fmax: int = None, + num_quantizers: int = 20, + codebook_size: list[int] = None, + # [1024,1024,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128] + threshold_ema_dead_code: int = 2, + position_embedding_type: str = "rope", + rope_theta: int = 10000, + rope_type: str = "default", + ln_type: str = "LayerNorm", + vocoder_attention_heads: int = 16, + vocoder_attn_window_size: list[int] = None, # [40,10] + use_sharding: bool = False, + shd_cfg: MiMoShardingCfg | None = None, + **kwargs, ): super().__init__(**kwargs) self.max_audio_seconds = max_audio_seconds @@ -165,22 +160,14 @@ def __init__( self.encoder_attention_heads = encoder_attention_heads self.encoder_ffn_dim = encoder_ffn_dim self.encoder_causal = encoder_causal - self.encoder_attn_window_size = ( - encoder_attn_window_size - if encoder_attn_window_size is not None - else [-1, -1] - ) + self.encoder_attn_window_size = encoder_attn_window_size if encoder_attn_window_size is not None else [-1, -1] self.decoder_layers = decoder_layers self.decoder_attention_heads = decoder_attention_heads self.decoder_ffn_dim = decoder_ffn_dim self.decoder_kernel_size = decoder_kernel_size self.decoder_stride_size = decoder_stride_size self.decoder_causal = decoder_causal - self.decoder_attn_window_size = ( - decoder_attn_window_size - if decoder_attn_window_size is not None - else [-1, -1] - ) + self.decoder_attn_window_size = decoder_attn_window_size if decoder_attn_window_size is not None else [-1, -1] self.nfft = nfft self.vocoder_dim = vocoder_dim self.vocoder_intermediate_dim = vocoder_intermediate_dim @@ -200,11 +187,7 @@ def __init__( self.rope_type = rope_type self.ln_type = ln_type self.vocoder_attention_heads = vocoder_attention_heads - self.vocoder_attn_window_size = ( - vocoder_attn_window_size - if vocoder_attn_window_size is not None - else [40, 10] - ) + self.vocoder_attn_window_size = vocoder_attn_window_size if vocoder_attn_window_size is not None else [40, 10] # Sharding configuration if shd_cfg is None: diff --git a/bonsai/models/mimo_audio/mimo_audio_tokenizer_params.py b/bonsai/models/mimo_audio/mimo_audio_tokenizer_params.py index 9035da73..0af1076f 100644 --- a/bonsai/models/mimo_audio/mimo_audio_tokenizer_params.py +++ b/bonsai/models/mimo_audio/mimo_audio_tokenizer_params.py @@ -37,18 +37,54 @@ def _get_key_mapping(config: model_lib.MiMoAudioTokenizerConfig) -> dict[str, tu for idx in range(config.encoder_layers): layer_mappings = { - rf"encoder\.layers\.{idx}\.self_attn\.q_proj\.weight": (f"encoder.layers.{idx}.self_attn.q_proj.kernel", TRANSFORM_LINEAR), - rf"encoder\.layers\.{idx}\.self_attn\.q_proj\.bias": (f"encoder.layers.{idx}.self_attn.q_proj.bias", TRANSFORM_NONE), - rf"encoder\.layers\.{idx}\.self_attn\.k_proj\.weight": (f"encoder.layers.{idx}.self_attn.k_proj.kernel", TRANSFORM_LINEAR), - rf"encoder\.layers\.{idx}\.self_attn\.k_proj\.bias": (f"encoder.layers.{idx}.self_attn.k_proj.bias", TRANSFORM_NONE), - rf"encoder\.layers\.{idx}\.self_attn\.v_proj\.weight": (f"encoder.layers.{idx}.self_attn.v_proj.kernel", TRANSFORM_LINEAR), - rf"encoder\.layers\.{idx}\.self_attn\.v_proj\.bias": (f"encoder.layers.{idx}.self_attn.v_proj.bias", TRANSFORM_NONE), - rf"encoder\.layers\.{idx}\.self_attn\.out_proj\.weight": (f"encoder.layers.{idx}.self_attn.out_proj.kernel", TRANSFORM_LINEAR), - rf"encoder\.layers\.{idx}\.self_attn\.out_proj\.bias": (f"encoder.layers.{idx}.self_attn.out_proj.bias", TRANSFORM_NONE), - rf"encoder\.layers\.{idx}\.self_attn_layer_norm\.weight": (f"encoder.layers.{idx}.self_attn_layer_norm.scale", TRANSFORM_NONE), - rf"encoder\.layers\.{idx}\.self_attn_layer_norm\.bias": (f"encoder.layers.{idx}.self_attn_layer_norm.bias", TRANSFORM_NONE), - rf"encoder\.layers\.{idx}\.final_layer_norm\.weight": (f"encoder.layers.{idx}.final_layer_norm.scale", TRANSFORM_NONE), - rf"encoder\.layers\.{idx}\.final_layer_norm\.bias": (f"encoder.layers.{idx}.final_layer_norm.bias", TRANSFORM_NONE), + rf"encoder\.layers\.{idx}\.self_attn\.q_proj\.weight": ( + f"encoder.layers.{idx}.self_attn.q_proj.kernel", + TRANSFORM_LINEAR, + ), + rf"encoder\.layers\.{idx}\.self_attn\.q_proj\.bias": ( + f"encoder.layers.{idx}.self_attn.q_proj.bias", + TRANSFORM_NONE, + ), + rf"encoder\.layers\.{idx}\.self_attn\.k_proj\.weight": ( + f"encoder.layers.{idx}.self_attn.k_proj.kernel", + TRANSFORM_LINEAR, + ), + rf"encoder\.layers\.{idx}\.self_attn\.k_proj\.bias": ( + f"encoder.layers.{idx}.self_attn.k_proj.bias", + TRANSFORM_NONE, + ), + rf"encoder\.layers\.{idx}\.self_attn\.v_proj\.weight": ( + f"encoder.layers.{idx}.self_attn.v_proj.kernel", + TRANSFORM_LINEAR, + ), + rf"encoder\.layers\.{idx}\.self_attn\.v_proj\.bias": ( + f"encoder.layers.{idx}.self_attn.v_proj.bias", + TRANSFORM_NONE, + ), + rf"encoder\.layers\.{idx}\.self_attn\.out_proj\.weight": ( + f"encoder.layers.{idx}.self_attn.out_proj.kernel", + TRANSFORM_LINEAR, + ), + rf"encoder\.layers\.{idx}\.self_attn\.out_proj\.bias": ( + f"encoder.layers.{idx}.self_attn.out_proj.bias", + TRANSFORM_NONE, + ), + rf"encoder\.layers\.{idx}\.self_attn_layer_norm\.weight": ( + f"encoder.layers.{idx}.self_attn_layer_norm.scale", + TRANSFORM_NONE, + ), + rf"encoder\.layers\.{idx}\.self_attn_layer_norm\.bias": ( + f"encoder.layers.{idx}.self_attn_layer_norm.bias", + TRANSFORM_NONE, + ), + rf"encoder\.layers\.{idx}\.final_layer_norm\.weight": ( + f"encoder.layers.{idx}.final_layer_norm.scale", + TRANSFORM_NONE, + ), + rf"encoder\.layers\.{idx}\.final_layer_norm\.bias": ( + f"encoder.layers.{idx}.final_layer_norm.bias", + TRANSFORM_NONE, + ), rf"encoder\.layers\.{idx}\.fc1\.weight": (f"encoder.layers.{idx}.fc1.kernel", TRANSFORM_LINEAR), rf"encoder\.layers\.{idx}\.fc1\.bias": (f"encoder.layers.{idx}.fc1.bias", TRANSFORM_NONE), rf"encoder\.layers\.{idx}\.fc2\.weight": (f"encoder.layers.{idx}.fc2.kernel", TRANSFORM_LINEAR), @@ -56,12 +92,14 @@ def _get_key_mapping(config: model_lib.MiMoAudioTokenizerConfig) -> dict[str, tu } mapping.update(layer_mappings) - mapping.update({ - r"encoder\.down_sample_layer\.0\.weight": ("encoder.down_sample_layer.kernel", TRANSFORM_CONV1D), - r"encoder\.down_sample_layer\.0\.bias": ("encoder.down_sample_layer.bias", TRANSFORM_NONE), - r"encoder\.down_sample_norm\.weight": ("encoder.down_norm.scale", TRANSFORM_NONE), - r"encoder\.down_sample_norm\.bias": ("encoder.down_norm.bias", TRANSFORM_NONE), - }) + mapping.update( + { + r"encoder\.down_sample_layer\.0\.weight": ("encoder.down_sample_layer.kernel", TRANSFORM_CONV1D), + r"encoder\.down_sample_layer\.0\.bias": ("encoder.down_sample_layer.bias", TRANSFORM_NONE), + r"encoder\.down_sample_norm\.weight": ("encoder.down_norm.scale", TRANSFORM_NONE), + r"encoder\.down_sample_norm\.bias": ("encoder.down_norm.bias", TRANSFORM_NONE), + } + ) for idx in range(config.num_quantizers): mapping[rf"encoder\.quantizer\.vq\.layers\.{idx}\._codebook\.embed"] = ( @@ -69,33 +107,71 @@ def _get_key_mapping(config: model_lib.MiMoAudioTokenizerConfig) -> dict[str, tu TRANSFORM_NONE, ) - mapping.update({ - r"decoder\.dconv1\.conv\.weight": ("decoder.dconv1.conv.kernel", TRANSFORM_NONE), - r"decoder\.dconv1\.conv\.bias": ("decoder.dconv1.conv.bias", TRANSFORM_NONE), - r"decoder\.dconv1\.norm\.weight": ("decoder.dconv1.norm.scale", TRANSFORM_NONE), - r"decoder\.dconv1\.norm\.bias": ("decoder.dconv1.norm.bias", TRANSFORM_NONE), - r"decoder\.layer_norm\.weight": ("decoder.layer_norm.scale", TRANSFORM_NONE), - r"decoder\.layer_norm\.bias": ("decoder.layer_norm.bias", TRANSFORM_NONE), - r"decoder\.dconv2\.conv\.weight": ("decoder.dconv2.conv.kernel", TRANSFORM_NONE), - r"decoder\.dconv2\.conv\.bias": ("decoder.dconv2.conv.bias", TRANSFORM_NONE), - r"decoder\.dconv2\.norm\.weight": ("decoder.dconv2.norm.scale", TRANSFORM_NONE), - r"decoder\.dconv2\.norm\.bias": ("decoder.dconv2.norm.bias", TRANSFORM_NONE), - }) + mapping.update( + { + r"decoder\.dconv1\.conv\.weight": ("decoder.dconv1.conv.kernel", TRANSFORM_NONE), + r"decoder\.dconv1\.conv\.bias": ("decoder.dconv1.conv.bias", TRANSFORM_NONE), + r"decoder\.dconv1\.norm\.weight": ("decoder.dconv1.norm.scale", TRANSFORM_NONE), + r"decoder\.dconv1\.norm\.bias": ("decoder.dconv1.norm.bias", TRANSFORM_NONE), + r"decoder\.layer_norm\.weight": ("decoder.layer_norm.scale", TRANSFORM_NONE), + r"decoder\.layer_norm\.bias": ("decoder.layer_norm.bias", TRANSFORM_NONE), + r"decoder\.dconv2\.conv\.weight": ("decoder.dconv2.conv.kernel", TRANSFORM_NONE), + r"decoder\.dconv2\.conv\.bias": ("decoder.dconv2.conv.bias", TRANSFORM_NONE), + r"decoder\.dconv2\.norm\.weight": ("decoder.dconv2.norm.scale", TRANSFORM_NONE), + r"decoder\.dconv2\.norm\.bias": ("decoder.dconv2.norm.bias", TRANSFORM_NONE), + } + ) for idx in range(config.decoder_layers): layer_mappings = { - rf"decoder\.layers\.{idx}\.self_attn\.q_proj\.weight": (f"decoder.layers.{idx}.self_attn.q_proj.kernel", TRANSFORM_LINEAR), - rf"decoder\.layers\.{idx}\.self_attn\.q_proj\.bias": (f"decoder.layers.{idx}.self_attn.q_proj.bias", TRANSFORM_NONE), - rf"decoder\.layers\.{idx}\.self_attn\.k_proj\.weight": (f"decoder.layers.{idx}.self_attn.k_proj.kernel", TRANSFORM_LINEAR), - rf"decoder\.layers\.{idx}\.self_attn\.k_proj\.bias": (f"decoder.layers.{idx}.self_attn.k_proj.bias", TRANSFORM_NONE), - rf"decoder\.layers\.{idx}\.self_attn\.v_proj\.weight": (f"decoder.layers.{idx}.self_attn.v_proj.kernel", TRANSFORM_LINEAR), - rf"decoder\.layers\.{idx}\.self_attn\.v_proj\.bias": (f"decoder.layers.{idx}.self_attn.v_proj.bias", TRANSFORM_NONE), - rf"decoder\.layers\.{idx}\.self_attn\.out_proj\.weight": (f"decoder.layers.{idx}.self_attn.out_proj.kernel", TRANSFORM_LINEAR), - rf"decoder\.layers\.{idx}\.self_attn\.out_proj\.bias": (f"decoder.layers.{idx}.self_attn.out_proj.bias", TRANSFORM_NONE), - rf"decoder\.layers\.{idx}\.self_attn_layer_norm\.weight": (f"decoder.layers.{idx}.self_attn_layer_norm.scale", TRANSFORM_NONE), - rf"decoder\.layers\.{idx}\.self_attn_layer_norm\.bias": (f"decoder.layers.{idx}.self_attn_layer_norm.bias", TRANSFORM_NONE), - rf"decoder\.layers\.{idx}\.final_layer_norm\.weight": (f"decoder.layers.{idx}.final_layer_norm.scale", TRANSFORM_NONE), - rf"decoder\.layers\.{idx}\.final_layer_norm\.bias": (f"decoder.layers.{idx}.final_layer_norm.bias", TRANSFORM_NONE), + rf"decoder\.layers\.{idx}\.self_attn\.q_proj\.weight": ( + f"decoder.layers.{idx}.self_attn.q_proj.kernel", + TRANSFORM_LINEAR, + ), + rf"decoder\.layers\.{idx}\.self_attn\.q_proj\.bias": ( + f"decoder.layers.{idx}.self_attn.q_proj.bias", + TRANSFORM_NONE, + ), + rf"decoder\.layers\.{idx}\.self_attn\.k_proj\.weight": ( + f"decoder.layers.{idx}.self_attn.k_proj.kernel", + TRANSFORM_LINEAR, + ), + rf"decoder\.layers\.{idx}\.self_attn\.k_proj\.bias": ( + f"decoder.layers.{idx}.self_attn.k_proj.bias", + TRANSFORM_NONE, + ), + rf"decoder\.layers\.{idx}\.self_attn\.v_proj\.weight": ( + f"decoder.layers.{idx}.self_attn.v_proj.kernel", + TRANSFORM_LINEAR, + ), + rf"decoder\.layers\.{idx}\.self_attn\.v_proj\.bias": ( + f"decoder.layers.{idx}.self_attn.v_proj.bias", + TRANSFORM_NONE, + ), + rf"decoder\.layers\.{idx}\.self_attn\.out_proj\.weight": ( + f"decoder.layers.{idx}.self_attn.out_proj.kernel", + TRANSFORM_LINEAR, + ), + rf"decoder\.layers\.{idx}\.self_attn\.out_proj\.bias": ( + f"decoder.layers.{idx}.self_attn.out_proj.bias", + TRANSFORM_NONE, + ), + rf"decoder\.layers\.{idx}\.self_attn_layer_norm\.weight": ( + f"decoder.layers.{idx}.self_attn_layer_norm.scale", + TRANSFORM_NONE, + ), + rf"decoder\.layers\.{idx}\.self_attn_layer_norm\.bias": ( + f"decoder.layers.{idx}.self_attn_layer_norm.bias", + TRANSFORM_NONE, + ), + rf"decoder\.layers\.{idx}\.final_layer_norm\.weight": ( + f"decoder.layers.{idx}.final_layer_norm.scale", + TRANSFORM_NONE, + ), + rf"decoder\.layers\.{idx}\.final_layer_norm\.bias": ( + f"decoder.layers.{idx}.final_layer_norm.bias", + TRANSFORM_NONE, + ), rf"decoder\.layers\.{idx}\.fc1\.weight": (f"decoder.layers.{idx}.fc1.kernel", TRANSFORM_LINEAR), rf"decoder\.layers\.{idx}\.fc1\.bias": (f"decoder.layers.{idx}.fc1.bias", TRANSFORM_NONE), rf"decoder\.layers\.{idx}\.fc2\.weight": (f"decoder.layers.{idx}.fc2.kernel", TRANSFORM_LINEAR), @@ -103,33 +179,77 @@ def _get_key_mapping(config: model_lib.MiMoAudioTokenizerConfig) -> dict[str, tu } mapping.update(layer_mappings) - mapping.update({ - r"decoder\.vocoder\.embeddings\.weight": ("decoder.vocoder.embeddings.kernel", TRANSFORM_LINEAR), - r"decoder\.vocoder\.embeddings\.bias": ("decoder.vocoder.embeddings.bias", TRANSFORM_NONE), - r"decoder\.vocoder\.layer_norm\.weight": ("decoder.vocoder.layer_norm.scale", TRANSFORM_NONE), - r"decoder\.vocoder\.layer_norm\.bias": ("decoder.vocoder.layer_norm.bias", TRANSFORM_NONE), - r"decoder\.vocoder\.head\.out\.weight": ("decoder.vocoder.head.linear.kernel", TRANSFORM_LINEAR), - r"decoder\.vocoder\.head\.out\.bias": ("decoder.vocoder.head.linear.bias", TRANSFORM_NONE), - r"decoder\.vocoder\.head\.istft\.window": ("decoder.vocoder.head.istft.window", TRANSFORM_NONE), - }) + mapping.update( + { + r"decoder\.vocoder\.embeddings\.weight": ("decoder.vocoder.embeddings.kernel", TRANSFORM_LINEAR), + r"decoder\.vocoder\.embeddings\.bias": ("decoder.vocoder.embeddings.bias", TRANSFORM_NONE), + r"decoder\.vocoder\.layer_norm\.weight": ("decoder.vocoder.layer_norm.scale", TRANSFORM_NONE), + r"decoder\.vocoder\.layer_norm\.bias": ("decoder.vocoder.layer_norm.bias", TRANSFORM_NONE), + r"decoder\.vocoder\.head\.out\.weight": ("decoder.vocoder.head.linear.kernel", TRANSFORM_LINEAR), + r"decoder\.vocoder\.head\.out\.bias": ("decoder.vocoder.head.linear.bias", TRANSFORM_NONE), + r"decoder\.vocoder\.head\.istft\.window": ("decoder.vocoder.head.istft.window", TRANSFORM_NONE), + } + ) for idx in range(config.vocoder_num_layers): layer_mappings = { - rf"decoder\.vocoder\.layers\.{idx}\.self_attn\.q_proj\.weight": (f"decoder.vocoder.layers.{idx}.self_attn.q_proj.kernel", TRANSFORM_LINEAR), - rf"decoder\.vocoder\.layers\.{idx}\.self_attn\.q_proj\.bias": (f"decoder.vocoder.layers.{idx}.self_attn.q_proj.bias", TRANSFORM_NONE), - rf"decoder\.vocoder\.layers\.{idx}\.self_attn\.k_proj\.weight": (f"decoder.vocoder.layers.{idx}.self_attn.k_proj.kernel", TRANSFORM_LINEAR), - rf"decoder\.vocoder\.layers\.{idx}\.self_attn\.k_proj\.bias": (f"decoder.vocoder.layers.{idx}.self_attn.k_proj.bias", TRANSFORM_NONE), - rf"decoder\.vocoder\.layers\.{idx}\.self_attn\.v_proj\.weight": (f"decoder.vocoder.layers.{idx}.self_attn.v_proj.kernel", TRANSFORM_LINEAR), - rf"decoder\.vocoder\.layers\.{idx}\.self_attn\.v_proj\.bias": (f"decoder.vocoder.layers.{idx}.self_attn.v_proj.bias", TRANSFORM_NONE), - rf"decoder\.vocoder\.layers\.{idx}\.self_attn\.out_proj\.weight": (f"decoder.vocoder.layers.{idx}.self_attn.out_proj.kernel", TRANSFORM_LINEAR), - rf"decoder\.vocoder\.layers\.{idx}\.self_attn\.out_proj\.bias": (f"decoder.vocoder.layers.{idx}.self_attn.out_proj.bias", TRANSFORM_NONE), - rf"decoder\.vocoder\.layers\.{idx}\.self_attn_layer_norm\.weight": (f"decoder.vocoder.layers.{idx}.self_attn_layer_norm.scale", TRANSFORM_NONE), - rf"decoder\.vocoder\.layers\.{idx}\.self_attn_layer_norm\.bias": (f"decoder.vocoder.layers.{idx}.self_attn_layer_norm.bias", TRANSFORM_NONE), - rf"decoder\.vocoder\.layers\.{idx}\.final_layer_norm\.weight": (f"decoder.vocoder.layers.{idx}.final_layer_norm.scale", TRANSFORM_NONE), - rf"decoder\.vocoder\.layers\.{idx}\.final_layer_norm\.bias": (f"decoder.vocoder.layers.{idx}.final_layer_norm.bias", TRANSFORM_NONE), - rf"decoder\.vocoder\.layers\.{idx}\.fc1\.weight": (f"decoder.vocoder.layers.{idx}.fc1.kernel", TRANSFORM_LINEAR), + rf"decoder\.vocoder\.layers\.{idx}\.self_attn\.q_proj\.weight": ( + f"decoder.vocoder.layers.{idx}.self_attn.q_proj.kernel", + TRANSFORM_LINEAR, + ), + rf"decoder\.vocoder\.layers\.{idx}\.self_attn\.q_proj\.bias": ( + f"decoder.vocoder.layers.{idx}.self_attn.q_proj.bias", + TRANSFORM_NONE, + ), + rf"decoder\.vocoder\.layers\.{idx}\.self_attn\.k_proj\.weight": ( + f"decoder.vocoder.layers.{idx}.self_attn.k_proj.kernel", + TRANSFORM_LINEAR, + ), + rf"decoder\.vocoder\.layers\.{idx}\.self_attn\.k_proj\.bias": ( + f"decoder.vocoder.layers.{idx}.self_attn.k_proj.bias", + TRANSFORM_NONE, + ), + rf"decoder\.vocoder\.layers\.{idx}\.self_attn\.v_proj\.weight": ( + f"decoder.vocoder.layers.{idx}.self_attn.v_proj.kernel", + TRANSFORM_LINEAR, + ), + rf"decoder\.vocoder\.layers\.{idx}\.self_attn\.v_proj\.bias": ( + f"decoder.vocoder.layers.{idx}.self_attn.v_proj.bias", + TRANSFORM_NONE, + ), + rf"decoder\.vocoder\.layers\.{idx}\.self_attn\.out_proj\.weight": ( + f"decoder.vocoder.layers.{idx}.self_attn.out_proj.kernel", + TRANSFORM_LINEAR, + ), + rf"decoder\.vocoder\.layers\.{idx}\.self_attn\.out_proj\.bias": ( + f"decoder.vocoder.layers.{idx}.self_attn.out_proj.bias", + TRANSFORM_NONE, + ), + rf"decoder\.vocoder\.layers\.{idx}\.self_attn_layer_norm\.weight": ( + f"decoder.vocoder.layers.{idx}.self_attn_layer_norm.scale", + TRANSFORM_NONE, + ), + rf"decoder\.vocoder\.layers\.{idx}\.self_attn_layer_norm\.bias": ( + f"decoder.vocoder.layers.{idx}.self_attn_layer_norm.bias", + TRANSFORM_NONE, + ), + rf"decoder\.vocoder\.layers\.{idx}\.final_layer_norm\.weight": ( + f"decoder.vocoder.layers.{idx}.final_layer_norm.scale", + TRANSFORM_NONE, + ), + rf"decoder\.vocoder\.layers\.{idx}\.final_layer_norm\.bias": ( + f"decoder.vocoder.layers.{idx}.final_layer_norm.bias", + TRANSFORM_NONE, + ), + rf"decoder\.vocoder\.layers\.{idx}\.fc1\.weight": ( + f"decoder.vocoder.layers.{idx}.fc1.kernel", + TRANSFORM_LINEAR, + ), rf"decoder\.vocoder\.layers\.{idx}\.fc1\.bias": (f"decoder.vocoder.layers.{idx}.fc1.bias", TRANSFORM_NONE), - rf"decoder\.vocoder\.layers\.{idx}\.fc2\.weight": (f"decoder.vocoder.layers.{idx}.fc2.kernel", TRANSFORM_LINEAR), + rf"decoder\.vocoder\.layers\.{idx}\.fc2\.weight": ( + f"decoder.vocoder.layers.{idx}.fc2.kernel", + TRANSFORM_LINEAR, + ), rf"decoder\.vocoder\.layers\.{idx}\.fc2\.bias": (f"decoder.vocoder.layers.{idx}.fc2.bias", TRANSFORM_NONE), } mapping.update(layer_mappings) @@ -137,9 +257,7 @@ def _get_key_mapping(config: model_lib.MiMoAudioTokenizerConfig) -> dict[str, tu return mapping -def _get_jax_key( - mapping: dict[str, tuple[str, Transform]], source_key: str -) -> tuple[str | None, Transform | None]: +def _get_jax_key(mapping: dict[str, tuple[str, Transform]], source_key: str) -> tuple[str | None, Transform | None]: for pat, (jax_key, transform) in mapping.items(): if re.fullmatch(pat, source_key): return jax_key, transform @@ -190,7 +308,6 @@ def load_tokenizer_weights_from_safetensors( mesh: jax.sharding.Mesh | None = None, rngs: nnx.Rngs | None = None, ) -> model_lib.FlaxMiMoAudioTokenizer: - model = nnx.eval_shape(lambda: model_lib.FlaxMiMoAudioTokenizer(config, dtype=dtype, rngs=nnx.Rngs(params=0))) graph_def, abs_state = nnx.split(model) state_dict = abs_state.to_pure_dict() @@ -227,17 +344,23 @@ def load_tokenizer_weights_from_safetensors( encoder_rotary_dim = config.d_model // config.encoder_attention_heads encoder_half_dim = encoder_rotary_dim // 2 - encoder_inv_freq = 1.0 / (config.rope_theta ** (jnp.arange(0, encoder_half_dim, dtype=jnp.float32) / float(encoder_half_dim))) + encoder_inv_freq = 1.0 / ( + config.rope_theta ** (jnp.arange(0, encoder_half_dim, dtype=jnp.float32) / float(encoder_half_dim)) + ) model.encoder.position_embedding.inv_freq.value = encoder_inv_freq decoder_rotary_dim = config.d_model // config.decoder_attention_heads decoder_half_dim = decoder_rotary_dim // 2 - decoder_inv_freq = 1.0 / (config.rope_theta ** (jnp.arange(0, decoder_half_dim, dtype=jnp.float32) / float(decoder_half_dim))) + decoder_inv_freq = 1.0 / ( + config.rope_theta ** (jnp.arange(0, decoder_half_dim, dtype=jnp.float32) / float(decoder_half_dim)) + ) model.decoder.position_embedding.inv_freq.value = decoder_inv_freq vocoder_rotary_dim = config.vocoder_dim // config.vocoder_attention_heads vocoder_half_dim = vocoder_rotary_dim // 2 - vocoder_inv_freq = 1.0 / (config.rope_theta ** (jnp.arange(0, vocoder_half_dim, dtype=jnp.float32) / float(vocoder_half_dim))) + vocoder_inv_freq = 1.0 / ( + config.rope_theta ** (jnp.arange(0, vocoder_half_dim, dtype=jnp.float32) / float(vocoder_half_dim)) + ) model.decoder.vocoder.position_embedding.inv_freq.value = vocoder_inv_freq window = jnp.hanning(config.nfft).astype(dtype) diff --git a/bonsai/models/mimo_audio/modeling.py b/bonsai/models/mimo_audio/modeling.py index 3d6196cc..877cbd27 100644 --- a/bonsai/models/mimo_audio/modeling.py +++ b/bonsai/models/mimo_audio/modeling.py @@ -2,14 +2,10 @@ import jax import jax.numpy as jnp from flax import nnx -from bonsai.models.qwen2.modeling import Qwen2, ModelConfig as Qwen2Config, Cache +from bonsai.models.qwen2.modeling import Qwen2, Cache from bonsai.models.qwen3.modeling import shard from bonsai.utils.samplers import Sampler, GreedySampler -from bonsai.models.mimo_audio.mimo_audio_configuration import ( - MiMoAudioConfig, - MiMoAudioArguments, - MiMoSamplerConfig -) +from bonsai.models.mimo_audio.mimo_audio_configuration import MiMoAudioConfig, MiMoAudioArguments, MiMoSamplerConfig class MiMoSampler: @@ -18,19 +14,12 @@ class MiMoSampler: def __init__(self, config: MiMoSamplerConfig): self.config = config if config.do_sample: - self._sampler = Sampler( - temperature=config.temperature, - top_k=config.top_k, - top_p=config.top_p - ) + self._sampler = Sampler(temperature=config.temperature, top_k=config.top_k, top_p=config.top_p) else: self._sampler = GreedySampler() def sample( - self, - logits: jnp.ndarray, - key: jax.random.PRNGKey, - removed_tokens: Optional[List[int]] = None + self, logits: jnp.ndarray, key: jax.random.PRNGKey, removed_tokens: Optional[List[int]] = None ) -> jnp.ndarray: if removed_tokens: for t in removed_tokens: @@ -42,11 +31,11 @@ def sample( class FlaxMiMoAudioForCausalLM(nnx.Module): def __init__( - self, - config: MiMoAudioConfig, - args: MiMoAudioArguments, - rngs: Optional[nnx.Rngs] = None, - dtype: jnp.dtype = jnp.bfloat16, + self, + config: MiMoAudioConfig, + args: MiMoAudioArguments, + rngs: Optional[nnx.Rngs] = None, + dtype: jnp.dtype = jnp.bfloat16, ): if rngs is None: rngs = nnx.Rngs(0) @@ -76,42 +65,31 @@ def __init__( self.input_local_transformer.embedder = None self.lm_head = shard( - nnx.Linear( - config.hidden_size, - config.vocab_size, - use_bias=False, - dtype=self.dtype, - rngs=rngs - ), - self.shd_cfg.emb_dv + nnx.Linear(config.hidden_size, config.vocab_size, use_bias=False, dtype=self.dtype, rngs=rngs), + self.shd_cfg.emb_dv, ) - self.local_transformer_lm_heads = nnx.List([ - shard( - nnx.Linear( - config.local_dim, - self.speech_vocab_sizes[i], - use_bias=False, - dtype=self.dtype, - rngs=rngs - ), - self.shd_cfg.emb_dv - ) - for i in range(self.audio_channels) - ]) - - self.speech_embeddings = nnx.List([ - shard( - nnx.Embed( - self.speech_vocab_sizes[i], - config.input_local_dim, - dtype=self.dtype, - rngs=rngs - ), - self.shd_cfg.emb_vd - ) - for i in range(self.audio_channels) - ]) + self.local_transformer_lm_heads = nnx.List( + [ + shard( + nnx.Linear( + config.local_dim, self.speech_vocab_sizes[i], use_bias=False, dtype=self.dtype, rngs=rngs + ), + self.shd_cfg.emb_dv, + ) + for i in range(self.audio_channels) + ] + ) + + self.speech_embeddings = nnx.List( + [ + shard( + nnx.Embed(self.speech_vocab_sizes[i], config.input_local_dim, dtype=self.dtype, rngs=rngs), + self.shd_cfg.emb_vd, + ) + for i in range(self.audio_channels) + ] + ) self.speech_group_downcast = shard( nnx.Linear( @@ -119,26 +97,18 @@ def __init__( config.hidden_size, use_bias=False, dtype=self.dtype, - rngs=rngs + rngs=rngs, ), - self.shd_cfg.ffw_weight_df + self.shd_cfg.ffw_weight_df, ) self.hidden_states_downcast = shard( - nnx.Linear( - config.hidden_size, - config.local_dim, - use_bias=False, - dtype=self.dtype, - rngs=rngs - ), - self.shd_cfg.ffw_weight_df + nnx.Linear(config.hidden_size, config.local_dim, use_bias=False, dtype=self.dtype, rngs=rngs), + self.shd_cfg.ffw_weight_df, ) def apply_input_local_transformer( - self, - speech_embeddings: jnp.ndarray, - cache: Optional[Cache] = None + self, speech_embeddings: jnp.ndarray, cache: Optional[Cache] = None ) -> jnp.ndarray: """Apply input local transformer to speech embeddings""" B, T_groups, group_size, hidden_size = speech_embeddings.shape @@ -148,11 +118,7 @@ def apply_input_local_transformer( if cache is None: cache = self.input_local_transformer.init_cache( - self.input_local_qwen2_config, - B * T_groups, - group_size, - generate_steps=0, - dtype=self.dtype + self.input_local_qwen2_config, B * T_groups, group_size, generate_steps=0, dtype=self.dtype ) x = input_embeddings @@ -162,24 +128,19 @@ def apply_input_local_transformer( return x.reshape(B, T_groups, group_size, hidden_size) - def _prepare_input_embeds( - self, - input_ids: jnp.ndarray, - text_embed_fn - ) -> jnp.ndarray: + def _prepare_input_embeds(self, input_ids: jnp.ndarray, text_embed_fn) -> jnp.ndarray: """Prepare input embeddings from interleaved text and speech tokens""" B = input_ids.shape[0] - text_input_ids = input_ids[:, 0, ::self.group_size] - speech_input_ids = input_ids[:, 1:, :].reshape( - B, self.audio_channels, -1, self.group_size - ).transpose(0, 2, 1, 3) + text_input_ids = input_ids[:, 0, :: self.group_size] + speech_input_ids = ( + input_ids[:, 1:, :].reshape(B, self.audio_channels, -1, self.group_size).transpose(0, 2, 1, 3) + ) is_speech = text_input_ids == self.args.empty_idx speech_embeds = jnp.zeros( - (B, is_speech.shape[1], self.group_size, self.config.input_local_dim), - dtype=self.dtype + (B, is_speech.shape[1], self.group_size, self.config.input_local_dim), dtype=self.dtype ) for idx in range(self.audio_channels): @@ -198,9 +159,7 @@ def _prepare_input_embeds( speech_embeds = speech_embeds * is_speech[:, :, None, None] T_groups = speech_embeds.shape[1] - speech_grouped_embeds = self.speech_group_downcast( - speech_embeds.reshape(B, T_groups, -1) - ) + speech_grouped_embeds = self.speech_group_downcast(speech_embeds.reshape(B, T_groups, -1)) text_input_ids_safe = jnp.where(text_input_ids == -100, 0, text_input_ids) text_embeds = text_embed_fn(text_input_ids_safe) @@ -212,13 +171,12 @@ def _prepare_input_embeds( return shard(output, self.shd_cfg.act_btd) def forward( - self, - input_ids: jnp.ndarray, - cache: Cache, - pad_id: int = 0, + self, + input_ids: jnp.ndarray, + cache: Cache, + pad_id: int = 0, ) -> Tuple[jnp.ndarray, jnp.ndarray, Cache]: - """Forward pass through the model""" - text_input_ids = input_ids[:, 0, ::self.group_size] + text_input_ids = input_ids[:, 0, :: self.group_size] def text_embed_fn(x): return self.model.embedder.embedding.value[x] @@ -228,26 +186,23 @@ def text_embed_fn(x): B, T_groups, _ = inputs_embeds.shape segment_ids = 1 * (text_input_ids != -100) - # Run through main transformer x = inputs_embeds for i, layer in enumerate(self.model.layers): x = layer(x, cache[i], segment_ids) - hidden_states = self.model.final_norm(x) # [B, T_groups, hidden_size] + hidden_states = self.model.final_norm(x) text_logits = self.lm_head(hidden_states[:, -1:, :]) # [B, 1, vocab_size] # Downcast hidden states for local transformer - local_hidden_states = self.hidden_states_downcast( - hidden_states[:, -1:, :] - ) # [B, 1, local_dim] + local_hidden_states = self.hidden_states_downcast(hidden_states[:, -1:, :]) # [B, 1, local_dim] return text_logits, local_hidden_states, cache def local_forward( - self, - local_embeds: jnp.ndarray, # [B, 1, local_dim] - key: jax.random.PRNGKey, - local_sampler: Optional[MiMoSampler] = None, + self, + local_embeds: jnp.ndarray, + key: jax.random.PRNGKey, + local_sampler: Optional[MiMoSampler] = None, ) -> jnp.ndarray: """ Generate audio tokens for one group using local transformer. @@ -263,10 +218,7 @@ def local_forward( B = local_embeds.shape[0] delay_iters = self.group_size + max(self.delay_pattern) - local_tokens = jnp.zeros( - (B, self.group_size, self.audio_channels), - dtype=jnp.int32 - ) + local_tokens = jnp.zeros((B, self.group_size, self.audio_channels), dtype=jnp.int32) if local_sampler is None: local_sampler = MiMoSampler(MiMoSamplerConfig()) @@ -282,9 +234,7 @@ def local_forward( segment_ids = jnp.ones((B, 1), dtype=jnp.int32) for t in range(delay_iters): - hidden_state, cache = _local_transformer_step_jit( - self.local_transformer, local_embeds, cache, segment_ids - ) + hidden_state, cache = _local_transformer_step_jit(self.local_transformer, local_embeds, cache, segment_ids) next_local_embeds = jnp.zeros_like(local_embeds) @@ -298,11 +248,7 @@ def local_forward( cur_logits = cur_lm_head(hidden_state[:, -1, :]) key, subkey = jax.random.split(key) - cur_token = local_sampler.sample( - cur_logits, - subkey, - removed_tokens=[cur_empty] - ) + cur_token = local_sampler.sample(cur_logits, subkey, removed_tokens=[cur_empty]) local_tokens = local_tokens.at[:, t - cur_start, idx].set(cur_token) @@ -315,10 +261,6 @@ def local_forward( return local_tokens -# ============================================================================ -# JIT-compiled functions for fast inference -# ============================================================================ - @jax.jit def _local_transformer_step_jit( local_transformer: nnx.Module, @@ -326,23 +268,6 @@ def _local_transformer_step_jit( cache: Cache, segment_ids: jnp.ndarray, ) -> Tuple[jnp.ndarray, Cache]: - """ - JIT-compiled single step of local transformer forward pass. - - This is a helper function to accelerate the inner loop of local_forward. - Being a module-level function (not instance method) allows JAX to properly - JIT compile it. - - Args: - local_transformer: The local transformer module - local_embeds: [B, 1, local_dim] - cache: Cache for local transformer - segment_ids: [B, 1] - - Returns: - hidden_state: [B, 1, local_dim] - cache: Updated cache (IMPORTANT for correct behavior) - """ x = local_embeds for i, layer in enumerate(local_transformer.layers): x = layer(x, cache[i], segment_ids) @@ -357,23 +282,6 @@ def forward_jit( cache: Cache, pad_id: int = 0, ) -> Tuple[jnp.ndarray, jnp.ndarray, Cache]: - """ - JIT-compiled forward pass for fast inference. - - Similar to qwen2's forward function, this returns the cache to enable - proper JAX tracing of stateful computations. - - Args: - model: FlaxMiMoAudioForCausalLM instance - input_ids: [B, audio_channels + 1, T * group_size] - cache: Cache for KV storage - pad_id: Padding token ID - - Returns: - text_logits: [B, 1, vocab_size] - local_hidden_states: [B, 1, local_dim] - cache: Updated cache (for JAX tracing) - """ text_logits, local_hidden_states, cache = model.forward(input_ids, cache, pad_id) return text_logits, local_hidden_states, cache @@ -384,42 +292,4 @@ def local_forward_jit( local_embeds: jnp.ndarray, key: jax.random.PRNGKey, ) -> jnp.ndarray: - """ - JIT-compiled local forward pass for audio generation. - - NOTE: This version uses greedy sampling (no sampler parameter for JIT simplicity). - For temperature-based sampling, use model.local_forward() directly. - - Args: - model: FlaxMiMoAudioForCausalLM instance - local_embeds: [B, 1, local_dim] - key: Random key (used if needed in future, currently greedy) - - Returns: - audio_tokens: [B, group_size, audio_channels] - """ - # Use greedy sampling for JIT-compiled version return model.local_forward(local_embeds, key, local_sampler=None) - - -# Example usage: -if __name__ == "__main__": - # Create configuration - config = MiMoAudioConfig() - args = MiMoAudioArguments( - model_name_or_path="mimo-audio", - sosp_idx=151646, - eosp_idx=151647, - sostm_idx=151648, - eostm_idx=151649, - eot_idx=151643, - empty_idx=151645, - ) - - # Create model - model = FlaxMiMoAudioForCausalLM(config,args) - - print("Model created successfully!") - print(f"Audio channels: {model.audio_channels}") - print(f"Group size: {model.group_size}") - print(f"Speech vocab sizes: {model.speech_vocab_sizes}") diff --git a/bonsai/models/mimo_audio/params.py b/bonsai/models/mimo_audio/params.py index 759707df..7cb2392b 100644 --- a/bonsai/models/mimo_audio/params.py +++ b/bonsai/models/mimo_audio/params.py @@ -14,18 +14,51 @@ def _get_qwen2_key_mapping(prefix: str) -> dict[str, tuple[str, Transform]]: return { rf"{prefix}\.embed_tokens\.weight": (f"{prefix}.embedder.embedding", TRANSFORM_NONE), - rf"{prefix}\.layers\.([0-9]+)\.self_attn\.q_proj\.weight": (rf"{prefix}.layers.\1.attn.q_proj.kernel", TRANSFORM_LINEAR), - rf"{prefix}\.layers\.([0-9]+)\.self_attn\.k_proj\.weight": (rf"{prefix}.layers.\1.attn.k_proj.kernel", TRANSFORM_LINEAR), - rf"{prefix}\.layers\.([0-9]+)\.self_attn\.v_proj\.weight": (rf"{prefix}.layers.\1.attn.v_proj.kernel", TRANSFORM_LINEAR), - rf"{prefix}\.layers\.([0-9]+)\.self_attn\.o_proj\.weight": (rf"{prefix}.layers.\1.attn.o_proj.kernel", TRANSFORM_LINEAR), - rf"{prefix}\.layers\.([0-9]+)\.self_attn\.q_proj\.bias": (rf"{prefix}.layers.\1.attn.q_proj.bias", q_bias_flatten), - rf"{prefix}\.layers\.([0-9]+)\.self_attn\.k_proj\.bias": (rf"{prefix}.layers.\1.attn.k_proj.bias", kv_bias_flatten), - rf"{prefix}\.layers\.([0-9]+)\.self_attn\.v_proj\.bias": (rf"{prefix}.layers.\1.attn.v_proj.bias", kv_bias_flatten), - rf"{prefix}\.layers\.([0-9]+)\.mlp\.gate_proj\.weight": (rf"{prefix}.layers.\1.mlp.gate_proj.kernel", TRANSFORM_LINEAR), - rf"{prefix}\.layers\.([0-9]+)\.mlp\.up_proj\.weight": (rf"{prefix}.layers.\1.mlp.up_proj.kernel", TRANSFORM_LINEAR), - rf"{prefix}\.layers\.([0-9]+)\.mlp\.down_proj\.weight": (rf"{prefix}.layers.\1.mlp.down_proj.kernel", TRANSFORM_LINEAR), + rf"{prefix}\.layers\.([0-9]+)\.self_attn\.q_proj\.weight": ( + rf"{prefix}.layers.\1.attn.q_proj.kernel", + TRANSFORM_LINEAR, + ), + rf"{prefix}\.layers\.([0-9]+)\.self_attn\.k_proj\.weight": ( + rf"{prefix}.layers.\1.attn.k_proj.kernel", + TRANSFORM_LINEAR, + ), + rf"{prefix}\.layers\.([0-9]+)\.self_attn\.v_proj\.weight": ( + rf"{prefix}.layers.\1.attn.v_proj.kernel", + TRANSFORM_LINEAR, + ), + rf"{prefix}\.layers\.([0-9]+)\.self_attn\.o_proj\.weight": ( + rf"{prefix}.layers.\1.attn.o_proj.kernel", + TRANSFORM_LINEAR, + ), + rf"{prefix}\.layers\.([0-9]+)\.self_attn\.q_proj\.bias": ( + rf"{prefix}.layers.\1.attn.q_proj.bias", + q_bias_flatten, + ), + rf"{prefix}\.layers\.([0-9]+)\.self_attn\.k_proj\.bias": ( + rf"{prefix}.layers.\1.attn.k_proj.bias", + kv_bias_flatten, + ), + rf"{prefix}\.layers\.([0-9]+)\.self_attn\.v_proj\.bias": ( + rf"{prefix}.layers.\1.attn.v_proj.bias", + kv_bias_flatten, + ), + rf"{prefix}\.layers\.([0-9]+)\.mlp\.gate_proj\.weight": ( + rf"{prefix}.layers.\1.mlp.gate_proj.kernel", + TRANSFORM_LINEAR, + ), + rf"{prefix}\.layers\.([0-9]+)\.mlp\.up_proj\.weight": ( + rf"{prefix}.layers.\1.mlp.up_proj.kernel", + TRANSFORM_LINEAR, + ), + rf"{prefix}\.layers\.([0-9]+)\.mlp\.down_proj\.weight": ( + rf"{prefix}.layers.\1.mlp.down_proj.kernel", + TRANSFORM_LINEAR, + ), rf"{prefix}\.norm\.weight": (f"{prefix}.final_norm.scale", TRANSFORM_NONE), - rf"{prefix}\.layers\.([0-9]+)\.input_layernorm\.weight": (rf"{prefix}.layers.\1.input_layernorm.scale", TRANSFORM_NONE), + rf"{prefix}\.layers\.([0-9]+)\.input_layernorm\.weight": ( + rf"{prefix}.layers.\1.input_layernorm.scale", + TRANSFORM_NONE, + ), rf"{prefix}\.layers\.([0-9]+)\.post_attention_layernorm\.weight": ( rf"{prefix}.layers.\1.post_attention_layernorm.scale", TRANSFORM_NONE, @@ -43,14 +76,15 @@ def _get_mimo_key_mapping(audio_channels: int) -> dict[str, tuple[str, Transform for i in range(audio_channels): mapping[rf"speech_embeddings\.{i}\.weight"] = (f"speech_embeddings.{i}.embedding", TRANSFORM_NONE) - mapping[rf"local_transformer_lm_heads\.{i}\.weight"] = (f"local_transformer_lm_heads.{i}.kernel", TRANSFORM_LINEAR) + mapping[rf"local_transformer_lm_heads\.{i}\.weight"] = ( + f"local_transformer_lm_heads.{i}.kernel", + TRANSFORM_LINEAR, + ) return mapping -def _get_jax_key( - mapping: dict[str, tuple[str, Transform]], source_key: str -) -> tuple[str | None, Transform | None]: +def _get_jax_key(mapping: dict[str, tuple[str, Transform]], source_key: str) -> tuple[str | None, Transform | None]: """Get JAX key from source key using regex mapping.""" for pat, (jax_key, transform) in mapping.items(): match = re.fullmatch(pat, source_key) diff --git a/bonsai/models/mimo_audio/test/run_model.py b/bonsai/models/mimo_audio/test/run_model.py index 6ac7e215..9d50970a 100644 --- a/bonsai/models/mimo_audio/test/run_model.py +++ b/bonsai/models/mimo_audio/test/run_model.py @@ -17,7 +17,7 @@ def load_audio_tokenizer(tokenizer_path: str): with open(config_path) as f: config_dict = json.load(f) - config_dict['use_sharding'] = False + config_dict["use_sharding"] = False config = MiMoAudioTokenizerConfig(**config_dict) safetensors_path = os.path.join(tokenizer_path, "model.safetensors") @@ -31,6 +31,7 @@ def load_audio_tokenizer(tokenizer_path: str): return tokenizer_model, config + def load_main_model(model_path: str): from bonsai.models.mimo_audio.mimo_audio_configuration import MiMoAudioConfig, MiMoAudioArguments from bonsai.models.mimo_audio.params import create_model_with_weights @@ -63,6 +64,7 @@ def load_main_model(model_path: str): return model, config, args, text_tokenizer + def insert_between(tokens: list, group_size: int, fill_value: int) -> list: if group_size <= 1: return tokens @@ -84,7 +86,7 @@ def run_inference( tokenizer_config, text_to_speak: str, max_steps: int = 100, - output_dir: str = "test_outputs" + output_dir: str = "test_outputs", ): from bonsai.models.mimo_audio.modeling import forward_jit, MiMoSampler from bonsai.models.mimo_audio.mimo_audio_configuration import MiMoSamplerConfig @@ -125,11 +127,8 @@ def run_inference( text_sampler = MiMoSampler(MiMoSamplerConfig(temperature=0.6, top_p=1.0, do_sample=True)) audio_sampler = MiMoSampler(MiMoSamplerConfig(temperature=0.9, top_p=0.95, do_sample=True)) - pad_id = text_tokenizer.pad_token_id - text_logits, local_hidden_states, cache = forward_jit( - main_model, input_ids, cache, pad_id - ) + text_logits, local_hidden_states, cache = forward_jit(main_model, input_ids, cache, pad_id) generated_text_tokens = [] generated_audio_tokens_list = [] @@ -157,11 +156,7 @@ def run_inference( generated_audio_tokens_list.append(audio_tokens_step) else: key, subkey = jax.random.split(key) - audio_tokens = main_model.local_forward( - local_hidden_states, - subkey, - audio_sampler - ) + audio_tokens = main_model.local_forward(local_hidden_states, subkey, audio_sampler) for t in range(group_size): audio_tokens_step = audio_tokens[0, t, :] @@ -184,14 +179,11 @@ def run_inference( for i in range(group_size): next_input = next_input.at[0, ch + 1, i].set(audio_tokens[0, i, ch]) - text_logits, local_hidden_states, cache = forward_jit( - main_model, next_input, cache, pad_id - ) + text_logits, local_hidden_states, cache = forward_jit(main_model, next_input, cache, pad_id) generated_text = text_tokenizer.decode(generated_text_tokens, skip_special_tokens=True) print(f"text token output: {generated_text}") - audio_tokens_array = jnp.stack(generated_audio_tokens_list, axis=0).T speech_empty_ids = main_model.speech_empty_ids @@ -202,12 +194,12 @@ def run_inference( not_empty = audio_tokens_array[ch, :] != empty_id is_real_audio_mask = is_real_audio_mask | not_empty - audio_tokens_array = audio_tokens_array[:, is_real_audio_mask] decoded_audio = tokenizer_model.decode(audio_tokens_array) os.makedirs(output_dir, exist_ok=True) import soundfile as sf + audio_path = os.path.join(output_dir, "generated_audio.wav") audio_np = np.array(decoded_audio[0, 0, :]) sample_rate = tokenizer_config.sampling_rate @@ -217,7 +209,6 @@ def run_inference( def main(): - model_name = "XiaomiMiMo/MiMo-Audio-7B-Instruct" tokenizer_name = "XiaomiMiMo/MiMo-Audio-Tokenizer" @@ -227,10 +218,12 @@ def main(): tokenizer_model, tokenizer_config = load_audio_tokenizer(tokenizer_path) main_model, config, args, text_tokenizer = load_main_model(model_path) - text_to_speak = ("And now here is my secret, a very simple secret:It is only with the heart that one can see rightly;" - "What is essential is invisible to the eye.It's the time you wasted for your rose that makes your rose so important." - "Men have forgotten this truth, but you must not forget it.You become responsible for what you have tamed." - "You are responsible for your rose...") + text_to_speak = ( + "And now here is my secret, a very simple secret:It is only with the heart that one can see rightly;" + "What is essential is invisible to the eye.It's the time you wasted for your rose that makes your rose so important." + "Men have forgotten this truth, but you must not forget it.You become responsible for what you have tamed." + "You are responsible for your rose..." + ) run_inference( main_model=main_model, tokenizer_model=tokenizer_model, diff --git a/bonsai/models/mimo_audio/test/test_outputs_mimo_audio.py b/bonsai/models/mimo_audio/test/test_outputs_mimo_audio.py index f8a09c28..eb852cca 100644 --- a/bonsai/models/mimo_audio/test/test_outputs_mimo_audio.py +++ b/bonsai/models/mimo_audio/test/test_outputs_mimo_audio.py @@ -41,7 +41,7 @@ def setUpClass(cls): config_dict = json.load(f) config_kwargs = {k: v for k, v in config_dict.items() if k in MiMoAudioConfig.__dataclass_fields__} - config_kwargs['shd_cfg'] = ShardingCfg.no_sharding() + config_kwargs["shd_cfg"] = ShardingCfg.no_sharding() cls.bonsai_config = MiMoAudioConfig(**config_kwargs) cls.args = MiMoAudioArguments( @@ -55,13 +55,21 @@ def setUpClass(cls): ) torch_config = TorchMiMoAudioConfig.from_pretrained(model_ckpt_path) - cls.torch_model = TorchMiMoAudio.from_pretrained( - model_ckpt_path, config=torch_config, args=asdict(cls.args), torch_dtype=torch.float32 - ).eval().cpu() + cls.torch_model = ( + TorchMiMoAudio.from_pretrained( + model_ckpt_path, config=torch_config, args=asdict(cls.args), torch_dtype=torch.float32 + ) + .eval() + .cpu() + ) cls.nnx_model = params.create_model_with_weights( - model_path=model_ckpt_path, config=cls.bonsai_config, args=cls.args, - rngs=nnx.Rngs(0), dtype=jnp.float32, mesh=None, + model_path=model_ckpt_path, + config=cls.bonsai_config, + args=cls.args, + rngs=nnx.Rngs(0), + dtype=jnp.float32, + mesh=None, ) cls.batch_size = 1 @@ -72,16 +80,22 @@ def setUpClass(cls): def _init_cache(self, batch_size, token_len): return self.nnx_model.model.init_cache( - cfg=self.nnx_model.qwen2_config, batch_size=batch_size, - token_len=token_len, generate_steps=0, dtype=jnp.float32 + cfg=self.nnx_model.qwen2_config, + batch_size=batch_size, + token_len=token_len, + generate_steps=0, + dtype=jnp.float32, ) def _compare(self, jy, ty): if ty.dim() == 2 and jy.ndim == 3: ty = ty.unsqueeze(0) torch.testing.assert_close( - torch.tensor(np.array(jy, dtype=np.float32)), ty, - rtol=self.tol, atol=self.tol, check_dtype=False, + torch.tensor(np.array(jy, dtype=np.float32)), + ty, + rtol=self.tol, + atol=self.tol, + check_dtype=False, ) def test_text_embedder(self): @@ -117,8 +131,11 @@ def test_main_decoder_layer(self): with torch.no_grad(): ty_output = self.torch_model.model.layers[0].to(torch.float32)( - tx, position_ids=position_ids, position_embeddings=position_embeddings, - past_key_value=DynamicCache(), cache_position=cache_position, + tx, + position_ids=position_ids, + position_embeddings=position_embeddings, + past_key_value=DynamicCache(), + cache_position=cache_position, ) ty = ty_output[0] if isinstance(ty_output, tuple) else ty_output @@ -128,9 +145,9 @@ def test_all_main_decoder_layers(self): cache = self._init_cache(self.batch_size, self.num_input_tokens) shape = (self.batch_size, self.num_input_tokens, self.bonsai_config.hidden_size) - for layer_idx, (nm, tm, nc) in enumerate(zip( - self.nnx_model.model.layers, self.torch_model.model.layers, cache - )): + for layer_idx, (nm, tm, nc) in enumerate( + zip(self.nnx_model.model.layers, self.torch_model.model.layers, cache) + ): jx = jax.random.normal(jax.random.key(layer_idx), shape=shape) tx = torch.tensor(np.array(jx, dtype=np.float32)) segment_ids = jnp.ones((self.batch_size, self.num_input_tokens)) @@ -142,8 +159,11 @@ def test_all_main_decoder_layers(self): with torch.no_grad(): ty_output = tm.to(torch.float32)( - tx, position_ids=position_ids, position_embeddings=position_embeddings, - past_key_value=DynamicCache(), cache_position=cache_position, + tx, + position_ids=position_ids, + position_embeddings=position_embeddings, + past_key_value=DynamicCache(), + cache_position=cache_position, ) ty = ty_output[0] if isinstance(ty_output, tuple) else ty_output @@ -171,14 +191,22 @@ def test_main_self_attn(self): position_ids = cache_position.unsqueeze(0) position_embeddings = self.torch_model.model.rotary_emb(tx, position_ids) attention_mask = create_causal_mask( - config=self.torch_model.config, input_embeds=tx, attention_mask=None, - cache_position=cache_position, past_key_values=DynamicCache(), position_ids=position_ids, + config=self.torch_model.config, + input_embeds=tx, + attention_mask=None, + cache_position=cache_position, + past_key_values=DynamicCache(), + position_ids=position_ids, ) with torch.no_grad(): ty = self.torch_model.model.layers[0].self_attn.to(torch.float32)( - tx, position_ids=position_ids, position_embeddings=position_embeddings, - attention_mask=attention_mask, past_key_value=DynamicCache(), cache_position=cache_position, + tx, + position_ids=position_ids, + position_embeddings=position_embeddings, + attention_mask=attention_mask, + past_key_value=DynamicCache(), + cache_position=cache_position, )[0] self._compare(jy, ty) @@ -234,8 +262,11 @@ def test_local_transformer_layer(self): tx = torch.tensor(np.array(jx, dtype=np.float32)) cache = self.nnx_model.local_transformer.init_cache( - cfg=self.nnx_model.local_qwen2_config, batch_size=self.batch_size, - token_len=self.num_input_tokens, generate_steps=0, dtype=jnp.float32 + cfg=self.nnx_model.local_qwen2_config, + batch_size=self.batch_size, + token_len=self.num_input_tokens, + generate_steps=0, + dtype=jnp.float32, ) segment_ids = jnp.ones((self.batch_size, self.num_input_tokens)) jy = self.nnx_model.local_transformer.layers[0](jx, cache[0], segment_ids) @@ -246,8 +277,11 @@ def test_local_transformer_layer(self): with torch.no_grad(): ty_output = self.torch_model.local_transformer.layers[0].to(torch.float32)( - tx, position_ids=position_ids, position_embeddings=position_embeddings, - past_key_value=DynamicCache(), cache_position=cache_position, + tx, + position_ids=position_ids, + position_embeddings=position_embeddings, + past_key_value=DynamicCache(), + cache_position=cache_position, ) ty = ty_output[0] if isinstance(ty_output, tuple) else ty_output @@ -259,8 +293,11 @@ def test_input_local_transformer_layer(self): tx = torch.tensor(np.array(jx, dtype=np.float32)) cache = self.nnx_model.input_local_transformer.init_cache( - cfg=self.nnx_model.input_local_qwen2_config, batch_size=self.batch_size, - token_len=self.num_input_tokens, generate_steps=0, dtype=jnp.float32 + cfg=self.nnx_model.input_local_qwen2_config, + batch_size=self.batch_size, + token_len=self.num_input_tokens, + generate_steps=0, + dtype=jnp.float32, ) segment_ids = jnp.ones((self.batch_size, self.num_input_tokens)) jy = self.nnx_model.input_local_transformer.layers[0](jx, cache[0], segment_ids) @@ -269,22 +306,29 @@ def test_input_local_transformer_layer(self): position_ids = cache_position.unsqueeze(0) position_embeddings = self.torch_model.input_local_transformer.rotary_emb(tx, position_ids) attention_mask = torch.ones( - (self.batch_size, 1, self.num_input_tokens, self.num_input_tokens), - dtype=torch.float32, device=tx.device + (self.batch_size, 1, self.num_input_tokens, self.num_input_tokens), dtype=torch.float32, device=tx.device ) with torch.no_grad(): ty_output = self.torch_model.input_local_transformer.layers[0].to(torch.float32)( - tx, position_ids=position_ids, position_embeddings=position_embeddings, - attention_mask=attention_mask, past_key_value=DynamicCache(), cache_position=cache_position, + tx, + position_ids=position_ids, + position_embeddings=position_embeddings, + attention_mask=attention_mask, + past_key_value=DynamicCache(), + cache_position=cache_position, ) ty = ty_output[0] if isinstance(ty_output, tuple) else ty_output self._compare(jy, ty) def test_apply_input_local_transformer(self): - shape = (self.batch_size, self.num_input_tokens // self.group_size, self.group_size, - self.bonsai_config.input_local_dim) + shape = ( + self.batch_size, + self.num_input_tokens // self.group_size, + self.group_size, + self.bonsai_config.input_local_dim, + ) jx = jax.random.normal(jax.random.key(0), shape=shape) tx = torch.tensor(np.array(jx, dtype=np.float32)) jy = self.nnx_model.apply_input_local_transformer(jx, cache=None) @@ -306,6 +350,7 @@ def test_prepare_input_embeds(self): def text_embed_fn_jax(x): return self.nnx_model.model.embedder.embedding.value[x] + jax_embeds = self.nnx_model._prepare_input_embeds(input_ids_jax, text_embed_fn_jax) with torch.no_grad(): @@ -333,8 +378,10 @@ def test_full_forward(self): position_ids = torch.arange(num_groups).unsqueeze(0).expand(self.batch_size, -1) cache_position = torch.arange(num_groups) outputs_torch = self.torch_model( - input_ids=input_ids_torch, attention_mask=attention_mask, - position_ids=position_ids, cache_position=cache_position, + input_ids=input_ids_torch, + attention_mask=attention_mask, + position_ids=position_ids, + cache_position=cache_position, ) text_logits_torch = outputs_torch.text_logits local_hidden_torch = outputs_torch.local_hidden_states diff --git a/bonsai/models/mimo_audio/test/test_outputs_mimo_audio_tokenizer.py b/bonsai/models/mimo_audio/test/test_outputs_mimo_audio_tokenizer.py index 4afca147..e35959ad 100644 --- a/bonsai/models/mimo_audio/test/test_outputs_mimo_audio_tokenizer.py +++ b/bonsai/models/mimo_audio/test/test_outputs_mimo_audio_tokenizer.py @@ -58,8 +58,7 @@ def move_to_cpu(module): config_dict = json.load(f) jax_config = JaxTokenizerConfig(**config_dict, use_sharding=False) cls.nnx_model = load_tokenizer_weights_from_safetensors( - config=jax_config, safetensors_path=safetensors_path, - dtype=jnp.float32, mesh=None, rngs=nnx.Rngs(params=0) + config=jax_config, safetensors_path=safetensors_path, dtype=jnp.float32, mesh=None, rngs=nnx.Rngs(params=0) ) cls.batch_size = 1 @@ -72,8 +71,11 @@ def _compare(self, jy, ty): if ty.dim() == 2 and jy.ndim == 3: ty = ty.unsqueeze(0) torch.testing.assert_close( - torch.tensor(np.array(jy, dtype=np.float32)), ty, - rtol=self.tol, atol=self.tol, check_dtype=False, + torch.tensor(np.array(jy, dtype=np.float32)), + ty, + rtol=self.tol, + atol=self.tol, + check_dtype=False, ) def test_encoder_conv1(self): @@ -172,9 +174,13 @@ def test_vocoder_istft_head(self): ty = self.torch_model.decoder.vocoder.head(x) torch.testing.assert_close( - torch.tensor(np.array(jy, dtype=np.float32)), ty, - rtol=1e-2, atol=1e-2, check_dtype=False, + torch.tensor(np.array(jy, dtype=np.float32)), + ty, + rtol=1e-2, + atol=1e-2, + check_dtype=False, ) + if __name__ == "__main__": absltest.main() diff --git a/bonsai/models/qwen2/modeling.py b/bonsai/models/qwen2/modeling.py index 7b5c71ff..f6449632 100644 --- a/bonsai/models/qwen2/modeling.py +++ b/bonsai/models/qwen2/modeling.py @@ -19,7 +19,7 @@ from flax import nnx from jax import P from jax import numpy as jnp -from jax.sharding import PartitionSpec, get_abstract_mesh +from jax.sharding import get_abstract_mesh from jaxtyping import Array from bonsai.models.qwen3.modeling import ( @@ -28,13 +28,11 @@ MLP, RMSNorm, ShardingCfg, - ShardingSpec, _generate_pos_embeddings, apply_rope, compute_positions_from_segment_ids, count_left_pads, count_right_pads, - reshard, shard, ) @@ -94,7 +92,6 @@ def qwen2_1_5b(cls, use_sharding: bool = False): # qwen2-1.5B tie_word_embeddings=True, ) - @classmethod def qwen2_7b(cls, use_sharding: bool = False): # qwen2-7B return cls._from_param( @@ -132,26 +129,19 @@ class Attention(nnx.Module): def __init__(self, cfg: ModelConfig, *, rngs: nnx.Rngs): self.shd_cfg = cfg.shd_cfg - # Standard Linear layers matching official Qwen2 implementation - # q_proj: [B, T, D] @ [D, N*H] -> [B, T, N*H] self.q_proj = shard( - nnx.Linear(cfg.emb_dim, cfg.num_heads * cfg.head_dim, use_bias=True, rngs=rngs), - self.shd_cfg.q_weight_ndh + nnx.Linear(cfg.emb_dim, cfg.num_heads * cfg.head_dim, use_bias=True, rngs=rngs), self.shd_cfg.q_weight_ndh ) - # k_proj: [B, T, D] @ [D, K*H] -> [B, T, K*H] self.k_proj = shard( nnx.Linear(cfg.emb_dim, cfg.num_kv_heads * cfg.head_dim, use_bias=True, rngs=rngs), - self.shd_cfg.kv_weight_ndh + self.shd_cfg.kv_weight_ndh, ) - # v_proj: [B, T, D] @ [D, K*H] -> [B, T, K*H] self.v_proj = shard( nnx.Linear(cfg.emb_dim, cfg.num_kv_heads * cfg.head_dim, use_bias=True, rngs=rngs), - self.shd_cfg.kv_weight_ndh + self.shd_cfg.kv_weight_ndh, ) - # o_proj: [B, T, N*H] @ [N*H, D] -> [B, T, D] self.o_proj = shard( - nnx.Linear(cfg.num_heads * cfg.head_dim, cfg.emb_dim, use_bias=False, rngs=rngs), - self.shd_cfg.o_weight_nhd + nnx.Linear(cfg.num_heads * cfg.head_dim, cfg.emb_dim, use_bias=False, rngs=rngs), self.shd_cfg.o_weight_nhd ) self.cfg = cfg @@ -161,7 +151,6 @@ def __init__(self, cfg: ModelConfig, *, rngs: nnx.Rngs): @jax.named_scope("attention") def __call__(self, x: Array, cache: LayerCache | None, segment_ids: Array) -> Array: - # Linear projections output [B, T, N*H] or [B, T, K*H], then reshape to [B, T, N/K, H] b, t = x.shape[:2] query_proj = self.q_proj(x).reshape(b, t, self.num_heads, self.head_dim) @@ -173,7 +162,6 @@ def __call__(self, x: Array, cache: LayerCache | None, segment_ids: Array) -> Ar value_proj = self.v_proj(x).reshape(b, t, self.num_kv_heads, self.head_dim) value_proj = shard(value_proj, self.shd_cfg.act_btnh) # [B, T, K, H] - # RoPE and Cache Logic left_pads = count_left_pads(segment_ids) left_pads = shard(left_pads, P(self.shd_cfg.act_btnh[0])) cache.start_ind.value = jnp.where(cache.start_ind.value < 0, left_pads, cache.start_ind.value) @@ -182,49 +170,39 @@ def __call__(self, x: Array, cache: LayerCache | None, segment_ids: Array) -> Ar query_proj = apply_rope(query_proj, sin, cos) key_proj = apply_rope(key_proj, sin, cos) - # Ensure dtype matches cache and preserve sharding - # astype can break sharding, so re-shard after dtype conversion cache_dtype = cache.k_cache.dtype value_proj = shard(value_proj.astype(cache_dtype), self.shd_cfg.act_btnh) key_proj = shard(key_proj.astype(cache_dtype), self.shd_cfg.act_btnh) - # Update K/V cache [B, S, K, H] slice_indices = (0, cache.cur_ind.value, 0, 0) cache.v_cache.value = jax.lax.dynamic_update_slice(cache.v_cache.value, value_proj, slice_indices) cache.k_cache.value = jax.lax.dynamic_update_slice(cache.k_cache.value, key_proj, slice_indices) b, t, n, h = query_proj.shape - # GQA reshape and attention logits query_proj_gqa = query_proj.reshape((b, t, self.num_kv_heads, self.n_rep, h)) attn_logits = jnp.einsum("BTKGH,BSKH->BTSKG", query_proj_gqa, cache.k_cache.value) * self.scale - # Masking and Softmax q_pos = cache.cur_ind.value + jnp.arange(t, dtype=jnp.int32)[None, :] - cache.start_ind.value[:, None] ts = jnp.arange(cache.size, dtype=jnp.int32) # (cache.size,) kv_segment_ids = (ts[None, :] >= cache.start_ind.value[:, None]) & (ts[None, :] < cache.cur_ind.value + t) k_pos = ts[None, :] - cache.start_ind.value[:, None] # (b, cache.size) - # Segment mask (always applied) segment_mask = kv_segment_ids[:, None, :] == segment_ids[:, :, None] - # Conditionally apply causal masking if self.cfg.use_causal_mask: causal_mask = k_pos[:, None, :] <= q_pos[:, :, None] - final_mask = causal_mask & segment_mask # (B, T, S) + final_mask = causal_mask & segment_mask else: - # Bidirectional attention: only use segment mask - final_mask = segment_mask # (B, T, S) + final_mask = segment_mask attn_mask = final_mask[:, :, :, None, None] attn_logits = jnp.where(attn_mask, attn_logits, _K_MASK) - # Softmax attn_weights = jax.nn.softmax(attn_logits.astype(jnp.float32), axis=2).astype(attn_logits.dtype) qkv = jnp.einsum("BTSKG,BSKH->BTKGH", attn_weights, cache.v_cache.value) qkv = qkv.reshape((b, t, n, h)) - # Reshape for o_proj: [B, T, N, H] -> [B, T, N*H] qkv_flat = qkv.reshape(b, t, n * h) output = self.o_proj(qkv_flat) @@ -268,11 +246,7 @@ def __init__(self, cfg: ModelConfig, *, rngs: nnx.Rngs): self.out_emb_shd = None if get_abstract_mesh().empty else cfg.shd_cfg.act_btd self.layers = nnx.List([DecoderLayer(cfg=cfg, rngs=rngs) for _ in range(cfg.num_layers)]) self.final_norm = RMSNorm(cfg.emb_dim, cfg, rngs=rngs) - # Standard Linear layer for lm_head - self.lm_head = shard( - nnx.Linear(cfg.emb_dim, cfg.vocab_size, use_bias=False, rngs=rngs), - cfg.shd_cfg.emb_dv - ) + self.lm_head = shard(nnx.Linear(cfg.emb_dim, cfg.vocab_size, use_bias=False, rngs=rngs), cfg.shd_cfg.emb_dv) def init_cache( self, cfg: ModelConfig, batch_size: int, token_len: int, generate_steps: int, dtype: jnp.dtype = jnp.bfloat16 @@ -286,10 +260,7 @@ def __call__(self, tokens, segment_ids, cache, num_right_pads): x = layer(x, cache[i], segment_ids) logits = self.lm_head(self.final_norm(x)) - # For generation/sampling, replicate all dimensions across devices - # This will trigger automatic all-gather to prepare for sampling if not get_abstract_mesh().empty: - # logits shape: [B, T, V], replicate all dims for sampling compatibility logits = shard(logits, P(None, None, None)) return logits @@ -301,4 +272,4 @@ def forward(model: nnx.Module, cache: Cache, tokens: Array, pad_id: int) -> tupl num_right_pads = count_right_pads(tokens, pad_id) logits = model(tokens, segment_ids, cache, num_right_pads) target_ind = tokens.shape[-1] - num_right_pads - 1 - return logits[:, target_ind], cache \ No newline at end of file + return logits[:, target_ind], cache diff --git a/bonsai/models/qwen2/params.py b/bonsai/models/qwen2/params.py index 5230f334..f474b187 100644 --- a/bonsai/models/qwen2/params.py +++ b/bonsai/models/qwen2/params.py @@ -23,31 +23,24 @@ class Transform: def _get_key_and_transform_mapping(cfg: model_lib.ModelConfig) -> dict[str, tuple[str | None, Transform | None]]: - # For Linear layers, we only need simple transpose from PyTorch's (out, in) to JAX's (in, out) - return { r"model\.embed_tokens\.weight": ("embedder.embedding", TRANSFORM_NONE), - # Attention projections: simple transpose for Linear layers r"model\.layers\.([0-9]+)\.self_attn\.q_proj\.weight": (r"layers.\1.attn.q_proj.kernel", TRANSFORM_LINEAR), r"model\.layers\.([0-9]+)\.self_attn\.k_proj\.weight": (r"layers.\1.attn.k_proj.kernel", TRANSFORM_LINEAR), r"model\.layers\.([0-9]+)\.self_attn\.v_proj\.weight": (r"layers.\1.attn.v_proj.kernel", TRANSFORM_LINEAR), r"model\.layers\.([0-9]+)\.self_attn\.o_proj\.weight": (r"layers.\1.attn.o_proj.kernel", TRANSFORM_LINEAR), - # Attention biases: no transformation needed r"model\.layers\.([0-9]+)\.self_attn\.q_proj\.bias": (r"layers.\1.attn.q_proj.bias", TRANSFORM_NONE), r"model\.layers\.([0-9]+)\.self_attn\.k_proj\.bias": (r"layers.\1.attn.k_proj.bias", TRANSFORM_NONE), r"model\.layers\.([0-9]+)\.self_attn\.v_proj\.bias": (r"layers.\1.attn.v_proj.bias", TRANSFORM_NONE), - # MLP projections r"model\.layers\.([0-9]+)\.mlp\.gate_proj\.weight": (r"layers.\1.mlp.gate_proj.kernel", TRANSFORM_LINEAR), r"model\.layers\.([0-9]+)\.mlp\.up_proj\.weight": (r"layers.\1.mlp.up_proj.kernel", TRANSFORM_LINEAR), r"model\.layers\.([0-9]+)\.mlp\.down_proj\.weight": (r"layers.\1.mlp.down_proj.kernel", TRANSFORM_LINEAR), - # Normalization layers r"model\.norm\.weight": ("final_norm.scale", TRANSFORM_NONE), r"model\.layers\.([0-9]+)\.input_layernorm\.weight": (r"layers.\1.input_layernorm.scale", TRANSFORM_NONE), r"model\.layers\.([0-9]+)\.post_attention_layernorm\.weight": ( r"layers.\1.post_attention_layernorm.scale", TRANSFORM_NONE, ), - # LM head r"lm_head\.weight": ("lm_head.kernel", TRANSFORM_LINEAR), } @@ -143,4 +136,4 @@ def create_model_from_safe_tensors( model = nnx.merge(graph_def, state_dict) gc.collect() - return model \ No newline at end of file + return model diff --git a/bonsai/models/qwen2/tests/__init__.py b/bonsai/models/qwen2/tests/__init__.py index 2aae938f..1337256a 100644 --- a/bonsai/models/qwen2/tests/__init__.py +++ b/bonsai/models/qwen2/tests/__init__.py @@ -10,4 +10,4 @@ # 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. \ No newline at end of file +# limitations under the License. diff --git a/bonsai/models/qwen2/tests/run_model.py b/bonsai/models/qwen2/tests/run_model.py index 89c5030b..198bd96a 100644 --- a/bonsai/models/qwen2/tests/run_model.py +++ b/bonsai/models/qwen2/tests/run_model.py @@ -7,15 +7,13 @@ from transformers import AutoTokenizer from bonsai.models.qwen2 import modeling, params -from bonsai.utils import GreedySampler, Sampler +from bonsai.utils import Sampler def tokenize(tokenizer, input: list[str], shd: P | None = None): pad_idx = tokenizer.pad_token_id lines = [ - tokenizer.apply_chat_template( - [{"role": "user", "content": l}], tokenize=False, add_generation_prompt=True - ) + tokenizer.apply_chat_template([{"role": "user", "content": l}], tokenize=False, add_generation_prompt=True) for l in input ] lines = [tokenizer.encode(line) for line in lines] @@ -24,7 +22,6 @@ def tokenize(tokenizer, input: list[str], shd: P | None = None): def run_model(): - model_name = "Qwen/Qwen2-7B" model_ckpt_path = snapshot_download(model_name) @@ -44,18 +41,12 @@ def run_model(): trust_remote_code=True, ) - print(f"Tokenizer special tokens:") - print(f" eos_token: {tokenizer.eos_token} (ID: {tokenizer.eos_token_id})") - print(f" pad_token: {tokenizer.pad_token} (ID: {tokenizer.pad_token_id})") - print() tokens = tokenize(tokenizer, query, batch_shd) batch_size, token_len = tokens.shape generate_steps = 1024 - print(f"\nLoading model...") model = params.create_model_from_safe_tensors(model_ckpt_path, config, mesh) - print(f"Model loaded successfully\n") cache = model.init_cache(config, batch_size, token_len, generate_steps) @@ -70,23 +61,20 @@ def run_model(): im_end_token_id = tokenizer.encode("<|im_end|>")[0] for i in range(generate_steps): - # CRITICAL: Split key for each step to avoid deterministic sampling key, subkey = jax.random.split(key) next_tokens = jit_sampler(logits, key=subkey) current_token_id = int(next_tokens.squeeze(-1)[0]) - # Only check for actual EOS token - is_eos = (next_tokens.squeeze(-1) == tokenizer.eos_token_id) - is_im_end = (next_tokens.squeeze(-1) == im_end_token_id) - + is_eos = next_tokens.squeeze(-1) == tokenizer.eos_token_id + is_im_end = next_tokens.squeeze(-1) == im_end_token_id finished = finished | is_eos | is_im_end tokens_list.append(next_tokens) if finished.all(): - print(f"✓ Generation stopped at step {i+1}/{generate_steps} (EOS token reached)") + print(f"✓ Generation stopped at step {i + 1}/{generate_steps} (EOS token reached)") break # Continue generation @@ -107,4 +95,4 @@ def run_model(): run_model() -__all__ = ["run_model"] \ No newline at end of file +__all__ = ["run_model"] diff --git a/bonsai/models/qwen2/tests/test_outputs_qwen2.py b/bonsai/models/qwen2/tests/test_outputs_qwen2.py index 219fbe5c..25293451 100644 --- a/bonsai/models/qwen2/tests/test_outputs_qwen2.py +++ b/bonsai/models/qwen2/tests/test_outputs_qwen2.py @@ -63,8 +63,7 @@ def _setup_torch_attn(self, input_embeddings: torch.Tensor, attention_mask: None attention_mask = torch.ones((batch_size, seq_length), dtype=torch.bool, device=input_embeddings.device) causal_mask = torch.triu( - torch.ones((seq_length, seq_length), dtype=torch.bool, device=input_embeddings.device), - diagonal=1 + torch.ones((seq_length, seq_length), dtype=torch.bool, device=input_embeddings.device), diagonal=1 ) causal_mask = causal_mask.unsqueeze(0).unsqueeze(0) causal_mask = causal_mask.expand(batch_size, 1, seq_length, seq_length) @@ -92,9 +91,7 @@ def _nnx_forward_logits(self, cache: modeling.Cache, tokens: jax.Array, dtype: D def _process_hf_tokens(self, query: list[str]): messages = [{"role": "user", "content": s} for s in query] - text = [ - self.tokenizer.apply_chat_template([m], tokenize=False, add_generation_prompt=True) for m in messages - ] + text = [self.tokenizer.apply_chat_template([m], tokenize=False, add_generation_prompt=True) for m in messages] model_inputs = self.tokenizer(text, return_tensors="pt", padding=True, padding_side="left").to( self.torch_model.device ) From 200d2178b0fa2b8037d9172f98b39426171f1585 Mon Sep 17 00:00:00 2001 From: Pratham Shah <113518804+coder0143@users.noreply.github.com> Date: Thu, 22 Jan 2026 05:35:42 +0530 Subject: [PATCH 14/18] Removing dinov3 model output class to facilitate jit compilation (#130) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * adding readme, modeling and params files for dinov3 * added run and testing * added readme and removed some basic imports * removed ignore statement for ide * Added colab link in markdown * changed config and made changes wrt it, added all outputs, removed cosine similarity in tests * adding more tests and some final changes * final changes to params * fixed ruff based text formatting * using random init torch model in test_outputs and removed run_model due to gated repo. * using hf_hub constants for path and pre-commit done for ruff * updating tests * rename test_outputs to be model specific * inital commit for qwen3vl * removed initial commit * vjepa2 base fm, classifier and params done * fixed params and testing for model porting * updating for additional testing and stability * more testing, ruff, classifier works well * fixed foundation model, reformatted modeling.py, fixed testing * adding readme * name in readme 😅 * removed unnecessary self.config = config * changes and fixes * Fixed resnet link in README (#127) Co-authored-by: James Chapman * Vjepa2 format fixes (#128) * vjepa2: Use opencv-python than torchcodec * Refactor test_outputs and forward in modeling.py * [CI] run_selective_tests: handle renamed paths (#129) * remove dinov3 output class to facilitate jit compilation --------- Co-authored-by: Jen Ha <25069493+jenriver@users.noreply.github.com> Co-authored-by: vfdev Co-authored-by: James Chapman --- bonsai/models/dinov3/modeling.py | 16 ++++++++-------- .../models/dinov3/tests/test_outputs_dinov3.py | 4 ++-- 2 files changed, 10 insertions(+), 10 deletions(-) diff --git a/bonsai/models/dinov3/modeling.py b/bonsai/models/dinov3/modeling.py index 5039573a..e550414e 100644 --- a/bonsai/models/dinov3/modeling.py +++ b/bonsai/models/dinov3/modeling.py @@ -1,17 +1,12 @@ import dataclasses from typing import Tuple +import jax import jax.numpy as jnp from flax import nnx from jax import Array -@dataclasses.dataclass -class Dinov3ViTModelOutput: - last_hidden_state: Array - pooler_output: Array - - @dataclasses.dataclass(frozen=True) class DINOv3ViTFlaxConfig: model_type = "dinov3_ViT" @@ -320,7 +315,7 @@ def __init__(self, config: DINOv3ViTFlaxConfig, rngs: nnx.Rngs): self.layer = nnx.List([Dinov3ViTLayer(config, rngs=rngs) for _ in range(config.num_hidden_layers)]) self.norm = nnx.LayerNorm(config.hidden_size, epsilon=config.layer_norm_eps, rngs=rngs) - def __call__(self, pixel_values: Array) -> Dinov3ViTModelOutput: + def __call__(self, pixel_values: Array): hidden_states = self.embeddings(pixel_values) position_embeddings = self.rope_embeddings(pixel_values) @@ -330,4 +325,9 @@ def __call__(self, pixel_values: Array) -> Dinov3ViTModelOutput: sequence_output = self.norm(hidden_states) pooled_output = sequence_output[:, 0, :] - return Dinov3ViTModelOutput(**{"last_hidden_state": sequence_output, "pooler_output": pooled_output}) + return {"last_hidden_state": sequence_output, "pooler_output": pooled_output} + + +@jax.jit() +def forward(model: Dinov3ViTModel, inputs: Array): + return model(inputs) diff --git a/bonsai/models/dinov3/tests/test_outputs_dinov3.py b/bonsai/models/dinov3/tests/test_outputs_dinov3.py index 3f4ee7e7..0a31a4ed 100644 --- a/bonsai/models/dinov3/tests/test_outputs_dinov3.py +++ b/bonsai/models/dinov3/tests/test_outputs_dinov3.py @@ -97,7 +97,7 @@ def test_last_hidden_state(self): with torch.inference_mode(): ty = self.baseline_model(tx).last_hidden_state - jy = self.bonsai_model(jx).last_hidden_state + jy = self.bonsai_model(jx)["last_hidden_state"] np_y = np.asarray(jax.device_get(jy)) ty_bonsai = torch.tensor(np_y, dtype=torch.float32) @@ -113,7 +113,7 @@ def test_pooled_output_embeddings(self): with torch.inference_mode(): ty = self.baseline_model(tx).pooler_output - jy = self.bonsai_model(jx).pooler_output + jy = self.bonsai_model(jx)["pooler_output"] np_y = np.asarray(jax.device_get(jy)) ty_bonsai = torch.tensor(np_y, dtype=torch.float32) From c9ecb68fa496560dbe0d716ff9793510c732449e Mon Sep 17 00:00:00 2001 From: vfdev Date: Thu, 22 Jan 2026 01:20:00 +0100 Subject: [PATCH 15/18] Fixed convnext tests and nnx warnings (#135) Co-authored-by: Jen Ha <25069493+jenriver@users.noreply.github.com> --- bonsai/models/convnext/modeling.py | 2 +- bonsai/models/convnext/params.py | 4 ++-- .../convnext/tests/test_outputs_ConvNext.py | 19 +++++++++++-------- bonsai/models/densenet121/params.py | 4 ++-- bonsai/models/dinov3/params.py | 4 ++-- bonsai/models/gemma3/params.py | 4 ++-- bonsai/models/llada_8b/params.py | 4 ++-- bonsai/models/qwen3/params.py | 4 ++-- bonsai/models/resnet/params.py | 4 ++-- bonsai/models/sam2/params.py | 4 ++-- bonsai/models/umt5/params.py | 2 +- bonsai/models/vae/params.py | 4 ++-- bonsai/models/vgg19/params.py | 4 ++-- bonsai/models/vit/params.py | 2 +- bonsai/models/vjepa2/params.py | 4 ++-- pyproject.toml | 1 + 16 files changed, 37 insertions(+), 33 deletions(-) diff --git a/bonsai/models/convnext/modeling.py b/bonsai/models/convnext/modeling.py index 8b1bbb3d..15408638 100644 --- a/bonsai/models/convnext/modeling.py +++ b/bonsai/models/convnext/modeling.py @@ -84,7 +84,7 @@ def __call__(self, x: jax.Array, *, rngs: jax.Array, train: bool): x = self.pwconv2(x) if self.gamma is not None: - x = self.gamma.value * x + x = self.gamma[...] * x return res + drop_path(x, self.drop_path_rate, rngs=rngs, train=train) diff --git a/bonsai/models/convnext/params.py b/bonsai/models/convnext/params.py index 9c3888fc..b3e06f33 100644 --- a/bonsai/models/convnext/params.py +++ b/bonsai/models/convnext/params.py @@ -154,7 +154,7 @@ def create_convnext_from_pretrained( model = model_lib.ConvNeXt(cfg=cfg, rngs=nnx.Rngs(params=0)) graph_def, abs_state = nnx.split(model) - jax_state = abs_state.to_pure_dict() + jax_state = nnx.to_pure_dict(abs_state) mapping = _get_key_and_transform_mapping() @@ -176,7 +176,7 @@ def create_convnext_from_pretrained( raise RuntimeError(f"Encountered {len(conversion_errors)} weight conversion errors. Log:\n{full_error_log}") if mesh is not None: - sharding = nnx.get_named_sharding(abs_state, mesh).to_pure_dict() + sharding = nnx.to_pure_dict(nnx.get_named_sharding(abs_state, mesh)) jax_state = jax.device_put(jax_state, sharding) else: jax_state = jax.device_put(jax_state, jax.devices()[0]) diff --git a/bonsai/models/convnext/tests/test_outputs_ConvNext.py b/bonsai/models/convnext/tests/test_outputs_ConvNext.py index 2d052842..c3150c77 100644 --- a/bonsai/models/convnext/tests/test_outputs_ConvNext.py +++ b/bonsai/models/convnext/tests/test_outputs_ConvNext.py @@ -1,5 +1,6 @@ import jax import jax.numpy as jnp +import numpy as np import torch from absl.testing import absltest, parameterized from huggingface_hub import snapshot_download @@ -31,13 +32,13 @@ def test_embeddings(self): nnx_emb = self.bonsai_model.embedding_layer jx = jax.random.normal(jax.random.key(0), self.image_shape, dtype=jnp.float32) - tx = torch.tensor(jx).permute(0, 3, 1, 2) + tx = torch.tensor(np.asarray(jx)).permute(0, 3, 1, 2) with torch.no_grad(): ty = torch_emb(tx) jy = nnx_emb(jx) - torch.testing.assert_close(torch.tensor(jy).permute(0, 3, 1, 2), ty, rtol=1e-5, atol=1e-5) + torch.testing.assert_close(torch.tensor(np.asarray(jy)).permute(0, 3, 1, 2), ty, rtol=1e-5, atol=1e-5) def test_blocks_isolated(self): jax_model = self.bonsai_model @@ -50,21 +51,23 @@ def test_blocks_isolated(self): for stage_idx, (stage_blocks, h, dim) in enumerate(zip(jax_model.stages, stage_img_sizes, stage_dims)): jx_input = jax.random.normal(key, (1, h, h, dim), dtype=jnp.float32) - tx_input = torch.tensor(jx_input).permute(0, 3, 1, 2) + tx_input = torch.tensor(np.asarray(jx_input)).permute(0, 3, 1, 2) for block_idx in range(len(stage_blocks.layers)): j_out = stage_blocks.layers[block_idx](jx_input, rngs=key, train=False) with torch.no_grad(): t_out = torch_model.encoder.stages[stage_idx].layers[block_idx](tx_input) - torch.testing.assert_close(torch.tensor(j_out), t_out.permute(0, 2, 3, 1), rtol=5e-4, atol=5e-4) + torch.testing.assert_close( + torch.tensor(np.asarray(j_out)), t_out.permute(0, 2, 3, 1), rtol=5e-4, atol=5e-4 + ) def test_full(self): jx = jax.random.normal(jax.random.key(0), self.image_shape, dtype=jnp.float32) - tx = torch.tensor(jx).permute(0, 3, 1, 2) + tx = torch.tensor(np.asarray(jx)).permute(0, 3, 1, 2) with torch.no_grad(): ty = self.baseline_model(tx).logits jy = self.bonsai_model(jx, rngs=jax.random.key(0)) - torch.testing.assert_close(torch.tensor(jy), ty) + torch.testing.assert_close(torch.tensor(np.asarray(jy)), ty) class TestModuleFullOtherConfigs(parameterized.TestCase): @@ -80,13 +83,13 @@ def test_full(self, model_size): image_shape = (2, 224, 224, 3) jx = jax.random.normal(jax.random.key(0), image_shape, dtype=jnp.float32) - tx = torch.tensor(jx).permute(0, 3, 1, 2) + tx = torch.tensor(np.asarray(jx)).permute(0, 3, 1, 2) with torch.no_grad(): ty = baseline_model(tx).logits jy = bonsai_model(jx, rngs=jax.random.key(0)) - torch.testing.assert_close(torch.tensor(jy), ty, rtol=2e-5, atol=2e-5) + torch.testing.assert_close(torch.tensor(np.asarray(jy)), ty, rtol=2e-5, atol=2e-5) if __name__ == "__main__": diff --git a/bonsai/models/densenet121/params.py b/bonsai/models/densenet121/params.py index b0797080..c957eb0d 100644 --- a/bonsai/models/densenet121/params.py +++ b/bonsai/models/densenet121/params.py @@ -156,7 +156,7 @@ def create_model_from_h5( densenet = nnx.eval_shape(lambda: model_lib.DenseNet(cfg, rngs=nnx.Rngs(params=0))) graph_def, abs_state = nnx.split(densenet) - state_dict = abs_state.to_pure_dict() + state_dict = nnx.to_pure_dict(abs_state) mapping = _get_key_and_transform_mapping(cfg) for st_key, tensor in tensor_dict.items(): @@ -167,7 +167,7 @@ def create_model_from_h5( _assign_weights(keys, tensor, state_dict, st_key, transform) if mesh is not None: - sharding = nnx.get_named_sharding(abs_state, mesh).to_pure_dict() + sharding = nnx.to_pure_dict(nnx.get_named_sharding(abs_state, mesh)) state_dict = jax.device_put(state_dict, sharding) else: state_dict = jax.device_put(state_dict, jax.devices()[0]) diff --git a/bonsai/models/dinov3/params.py b/bonsai/models/dinov3/params.py index 3c9c327c..5b482204 100644 --- a/bonsai/models/dinov3/params.py +++ b/bonsai/models/dinov3/params.py @@ -109,9 +109,9 @@ def create_model_from_safe_tensors( dinov3 = nnx.eval_shape(lambda: Dinov3ViTModel(cfg, rngs=nnx.Rngs(0))) graph_def, abs_state = nnx.split(dinov3) - state_dict = abs_state.to_pure_dict() + state_dict = nnx.to_pure_dict(abs_state) # Only use sharding if mesh is provided - sharding = nnx.get_named_sharding(abs_state, mesh).to_pure_dict() if mesh is not None else None + sharding = nnx.to_pure_dict(nnx.get_named_sharding(abs_state, mesh)) if mesh is not None else None key_mapping = _get_key_and_transform_mapping() diff --git a/bonsai/models/gemma3/params.py b/bonsai/models/gemma3/params.py index f6d6f5fb..b9ecd268 100644 --- a/bonsai/models/gemma3/params.py +++ b/bonsai/models/gemma3/params.py @@ -244,8 +244,8 @@ def create_gemma3_from_pretrained(file_dir: str, cfg: model_lib.ModelConfig, *, gemma3 = model_lib.Gemma3Model(cfg, rngs=nnx.Rngs(0)) graph_def, abs_state = nnx.split(gemma3) - jax_state = abs_state.to_pure_dict() - sharding = nnx.get_named_sharding(abs_state, mesh).to_pure_dict() if mesh is not None else None + jax_state = nnx.to_pure_dict(abs_state) + sharding = nnx.to_pure_dict(nnx.get_named_sharding(abs_state, mesh)) if mesh is not None else None mapping = _get_key_and_transform_mapping() for st_key, tensor in tensor_dict.items(): diff --git a/bonsai/models/llada_8b/params.py b/bonsai/models/llada_8b/params.py index 8d7532e2..0c7f93de 100644 --- a/bonsai/models/llada_8b/params.py +++ b/bonsai/models/llada_8b/params.py @@ -193,7 +193,7 @@ def create_llada_from_pretrained( # 2. Create uninitialized model model = nnx.eval_shape(lambda: model_lib.LLaDAModel(cfg=config, rngs=nnx.Rngs(params=0, dropout=0))) graph_def, abs_state = nnx.split(model) - jax_state = abs_state.to_pure_dict() + jax_state = nnx.to_pure_dict(abs_state) # 3. Assign known weights mapping = _get_key_and_transform_mapping() @@ -212,7 +212,7 @@ def create_llada_from_pretrained( # 5. Device placement if mesh is not None: - sharding = nnx.get_named_sharding(abs_state, mesh).to_pure_dict() + sharding = nnx.to_pure_dict(nnx.get_named_sharding(abs_state, mesh)) jax_state = jax.device_put(jax_state, sharding) else: jax_state = jax.device_put(jax_state, jax.devices()[0]) diff --git a/bonsai/models/qwen3/params.py b/bonsai/models/qwen3/params.py index d0dec3b3..8b7aa662 100644 --- a/bonsai/models/qwen3/params.py +++ b/bonsai/models/qwen3/params.py @@ -113,9 +113,9 @@ def create_model_from_safe_tensors( qwen3 = nnx.eval_shape(lambda: model_lib.Qwen3(cfg, rngs=nnx.Rngs(params=0))) graph_def, abs_state = nnx.split(qwen3) - state_dict = abs_state.to_pure_dict() + state_dict = nnx.to_pure_dict(abs_state) # Only use sharding if mesh is provided - sharding = nnx.get_named_sharding(abs_state, mesh).to_pure_dict() if mesh is not None else None + sharding = nnx.to_pure_dict(nnx.get_named_sharding(abs_state, mesh)) if mesh is not None else None key_mapping = _get_key_and_transform_mapping(cfg) conversion_errors = [] diff --git a/bonsai/models/resnet/params.py b/bonsai/models/resnet/params.py index b487ba31..87f8ba38 100644 --- a/bonsai/models/resnet/params.py +++ b/bonsai/models/resnet/params.py @@ -148,7 +148,7 @@ def create_resnet_from_pretrained( model = nnx.eval_shape(lambda: model_lib.ResNet(config, rngs=nnx.Rngs(params=0))) graph_def, abs_state = nnx.split(model) - jax_state = abs_state.to_pure_dict() + jax_state = nnx.to_pure_dict(abs_state) mapping = _get_key_and_transform_mapping(config) conversion_errors = [] @@ -169,7 +169,7 @@ def create_resnet_from_pretrained( raise RuntimeError(f"Encountered {len(conversion_errors)} weight conversion errors. Log:\n{full_error_log}") if mesh is not None: - sharding = nnx.get_named_sharding(abs_state, mesh).to_pure_dict() + sharding = nnx.to_pure_dict(nnx.get_named_sharding(abs_state, mesh)) jax_state = jax.device_put(jax_state, sharding) else: jax_state = jax.device_put(jax_state, jax.devices()[0]) diff --git a/bonsai/models/sam2/params.py b/bonsai/models/sam2/params.py index 5d799ea2..a881087e 100644 --- a/bonsai/models/sam2/params.py +++ b/bonsai/models/sam2/params.py @@ -712,7 +712,7 @@ def create_sam2_from_pretrained( # 2. Create uninitialized SAM2 nnx model sam2 = nnx.eval_shape(lambda: model_lib.build_sam2_model_from_config(config, rngs=nnx.Rngs(params=0, dropout=0))) graph_def, abs_state = nnx.split(sam2) - jax_state = abs_state.to_pure_dict() + jax_state = nnx.to_pure_dict(abs_state) # 3. Assign known weights mapping = _get_key_and_transform_mapping() @@ -731,7 +731,7 @@ def create_sam2_from_pretrained( # 5. Device placement if mesh is not None: - sharding = nnx.get_named_sharding(abs_state, mesh).to_pure_dict() + sharding = nnx.to_pure_dict(nnx.get_named_sharding(abs_state, mesh)) jax_state = jax.device_put(jax_state, sharding) else: jax_state = jax.device_put(jax_state, jax.devices()[0]) diff --git a/bonsai/models/umt5/params.py b/bonsai/models/umt5/params.py index 001759b5..ad80ce7e 100644 --- a/bonsai/models/umt5/params.py +++ b/bonsai/models/umt5/params.py @@ -314,7 +314,7 @@ def create_model( graph_def, abs_state = nnx.split(umt5) state_dict = nnx.to_pure_dict(abs_state) # Only use sharding if mesh is provided - sharding = nnx.get_named_sharding(abs_state, mesh).to_pure_dict() if mesh is not None else None + sharding = nnx.to_pure_dict(nnx.get_named_sharding(abs_state, mesh)) if mesh is not None else None if not key_mapping: key_mapping = _get_key_and_transform_mapping(cls, cfg) diff --git a/bonsai/models/vae/params.py b/bonsai/models/vae/params.py index 2348031b..99d52a9e 100644 --- a/bonsai/models/vae/params.py +++ b/bonsai/models/vae/params.py @@ -283,7 +283,7 @@ def create_model_from_safe_tensors( vae = nnx.eval_shape(lambda: model_lib.VAE(cfg=cfg, rngs=nnx.Rngs(params=0))) graph_def, abs_state = nnx.split(vae) - jax_state = abs_state.to_pure_dict() + jax_state = nnx.to_pure_dict(abs_state) mapping = _get_key_and_transform_mapping() @@ -295,7 +295,7 @@ def create_model_from_safe_tensors( _assign_weights(keys, tensor, jax_state, st_key, transform.value) if mesh is not None: - sharding = nnx.get_named_sharding(abs_state, mesh).to_pure_dict() + sharding = nnx.to_pure_dict(nnx.get_named_sharding(abs_state, mesh)) state_dict = jax.device_put(jax_state, sharding) else: state_dict = jax.device_put(jax_state, jax.devices()[0]) diff --git a/bonsai/models/vgg19/params.py b/bonsai/models/vgg19/params.py index f62e64a1..f8c38869 100644 --- a/bonsai/models/vgg19/params.py +++ b/bonsai/models/vgg19/params.py @@ -148,7 +148,7 @@ def create_model_from_h5( vgg = nnx.eval_shape(lambda: model_lib.VGG(cfg, rngs=nnx.Rngs(params=0))) graph_def, abs_state = nnx.split(vgg) - state_dict = abs_state.to_pure_dict() + state_dict = nnx.to_pure_dict(abs_state) mapping = _get_key_and_transform_mapping(cfg) conversion_errors = [] @@ -168,7 +168,7 @@ def create_model_from_h5( raise RuntimeError(f"Encountered {len(conversion_errors)} weight conversion errors. Log:\n{full_error_log}") if mesh is not None: - sharding = nnx.get_named_sharding(abs_state, mesh).to_pure_dict() + sharding = nnx.to_pure_dict(nnx.get_named_sharding(abs_state, mesh)) state_dict = jax.device_put(state_dict, sharding) else: state_dict = jax.device_put(state_dict, jax.devices()[0]) diff --git a/bonsai/models/vit/params.py b/bonsai/models/vit/params.py index 1742d98a..6a500e40 100644 --- a/bonsai/models/vit/params.py +++ b/bonsai/models/vit/params.py @@ -147,7 +147,7 @@ def create_vit_from_pretrained(file_dir: str, config: model_lib.ModelConfig): vit = model_lib.ViTClassificationModel(config, rngs=nnx.Rngs(0)) graph_def, abs_state = nnx.split(vit) - jax_state = abs_state.to_pure_dict() + jax_state = nnx.to_pure_dict(abs_state) mapping = _get_key_and_transform_mapping(config) conversion_errors = [] diff --git a/bonsai/models/vjepa2/params.py b/bonsai/models/vjepa2/params.py index e4b9ac37..c2f4213e 100644 --- a/bonsai/models/vjepa2/params.py +++ b/bonsai/models/vjepa2/params.py @@ -339,9 +339,9 @@ def create_model_from_safe_tensors( else: vjepa2 = nnx.eval_shape(lambda: VJEPA2Model(cfg, rngs=nnx.Rngs(0))) graph_def, abs_state = nnx.split(vjepa2) - state_dict = abs_state.to_pure_dict() + state_dict = nnx.to_pure_dict(abs_state) # Only use sharding if mesh is provided - sharding = nnx.get_named_sharding(abs_state, mesh).to_pure_dict() if mesh is not None else None + sharding = nnx.to_pure_dict(nnx.get_named_sharding(abs_state, mesh)) if mesh is not None else None key_mapping = _get_key_and_transform_mapping(classifier) diff --git a/pyproject.toml b/pyproject.toml index c1b8ee7b..720accb2 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -45,6 +45,7 @@ test-env = [ "timm", "h5py", "diffusers[flax]", + "keras_hub", ] testing = [ From 1f41221efb798c74a49ccde2c14632fee6547293 Mon Sep 17 00:00:00 2001 From: Moriyuki SUZUKI <139259020+Moriyuki-S@users.noreply.github.com> Date: Fri, 23 Jan 2026 03:07:37 +0900 Subject: [PATCH 16/18] fix: correct checkbox syntax in pull request template (#139) Change `- []` to `- [ ]` to ensure checkboxes are properly rendered by GitHub. --- .github/pull_request_template.md | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/.github/pull_request_template.md b/.github/pull_request_template.md index 875e290e..b466057b 100644 --- a/.github/pull_request_template.md +++ b/.github/pull_request_template.md @@ -8,7 +8,7 @@ Resolves #\ **Checklist** -- [] I have read the **[Contribution Guidelines](https://github.com/jax-ml/bonsai/blob/main/CONTRIBUTING.md#contributing-a-model)** and used [pre-commit hooks](https://github.com/jax-ml/bonsai/blob/main/CONTRIBUTING.md#linting-and-type-checking) to format this commit. -- [] I have added all the necessary **unit tests** for my change. (`run_model.py` for model usage, `test_outputs.py` and/or `model_validation_colab.ipynb` for quality). -- [] **(If using an LLM)** I have carefully reviewed and removed all **superfluous comments** or unneeded, commented-out code. Only necessary and functional code remains. -- [] I have signed the **[Contributor License Agreement (CLA)](https://cla.developers.google.com/about)**. \ No newline at end of file +- [ ] I have read the **[Contribution Guidelines](https://github.com/jax-ml/bonsai/blob/main/CONTRIBUTING.md#contributing-a-model)** and used [pre-commit hooks](https://github.com/jax-ml/bonsai/blob/main/CONTRIBUTING.md#linting-and-type-checking) to format this commit. +- [ ] I have added all the necessary **unit tests** for my change. (`run_model.py` for model usage, `test_outputs.py` and/or `model_validation_colab.ipynb` for quality). +- [ ] **(If using an LLM)** I have carefully reviewed and removed all **superfluous comments** or unneeded, commented-out code. Only necessary and functional code remains. +- [ ] I have signed the **[Contributor License Agreement (CLA)](https://cla.developers.google.com/about)**. From 4978280e0a37467c5184c6b9445b912a965dc709 Mon Sep 17 00:00:00 2001 From: Cosmo <150933444+CosmoNaught@users.noreply.github.com> Date: Thu, 22 Jan 2026 18:28:08 +0000 Subject: [PATCH 17/18] Mamba2: SSM State Caching for Efficient Incremental Generation (#131) * implement state space caching * refactor docstrings and comments * convert unicode string to ASCII-only * move create_empty_cache() to modeling.py --------- Co-authored-by: James Chapman --- bonsai/models/mamba2/modeling.py | 154 +++++++++++++----- bonsai/models/mamba2/tests/run_model.py | 148 +++++++++++++---- .../mamba2/tests/test_outputs_mamba_2.py | 101 +++++++++++- 3 files changed, 334 insertions(+), 69 deletions(-) diff --git a/bonsai/models/mamba2/modeling.py b/bonsai/models/mamba2/modeling.py index af2e6978..5184f5e1 100644 --- a/bonsai/models/mamba2/modeling.py +++ b/bonsai/models/mamba2/modeling.py @@ -79,6 +79,49 @@ def tiny(cls): return cls(vocab_size=1000, hidden_size=64, state_size=16, head_dim=16, chunk_size=32, num_hidden_layers=2) +@jax.tree_util.register_pytree_node_class +@dataclasses.dataclass +class Mamba2Cache: + """Cache for Mamba2 SSM and convolution states.""" + + ssm_states: list[jnp.ndarray] # (batch, heads, head_dim, state_size) per layer + conv_states: list[jnp.ndarray] # (batch, conv_dim, kernel_size - 1) per layer + + def tree_flatten(self): + return (self.ssm_states, self.conv_states), None + + @classmethod + def tree_unflatten(cls, aux_data, children): + return cls(ssm_states=list(children[0]), conv_states=list(children[1])) + + +def create_empty_cache( + cfg: Mamba2Config, + batch_size: int, + dtype: jnp.dtype = jnp.float32, +) -> Mamba2Cache: + """Create an empty cache for Mamba2 model. + + Args: + cfg: Mamba2Config for the model. + batch_size: Batch size for the cache. + dtype: Data type for cache arrays. + + Returns: + Empty Mamba2Cache with zero-initialized states. + """ + conv_dim = cfg.intermediate_size + 2 * cfg.state_size + cache_len = cfg.conv_kernel - 1 + + conv_states = [jnp.zeros((batch_size, conv_dim, cache_len), dtype=dtype) for _ in range(cfg.num_hidden_layers)] + ssm_states = [ + jnp.zeros((batch_size, cfg.num_heads, cfg.head_dim, cfg.state_size), dtype=dtype) + for _ in range(cfg.num_hidden_layers) + ] + + return Mamba2Cache(ssm_states=ssm_states, conv_states=conv_states) + + # SSD Core Algorithm @@ -228,7 +271,7 @@ def __call__(self, hidden_states: jnp.ndarray, residual: jnp.ndarray | None = No class DepthwiseConv1d(nnx.Module): - """Depthwise causal 1D convolution. Expects (batch, seq_len, channels).""" + """Depthwise causal 1D convolution with state caching. Expects (batch, seq_len, channels).""" def __init__(self, features: int, kernel_size: int, use_bias: bool = True, *, rngs: nnx.Rngs): self.features = features @@ -237,15 +280,25 @@ def __init__(self, features: int, kernel_size: int, use_bias: bool = True, *, rn in_features=features, out_features=features, kernel_size=(kernel_size,), - padding=((kernel_size - 1, 0),), + padding=((0, 0),), feature_group_count=features, use_bias=use_bias, rngs=rngs, ) @jax.named_scope("depthwise_conv1d") - def __call__(self, x: jnp.ndarray) -> jnp.ndarray: - return self.conv(x) + def __call__(self, x: jnp.ndarray, conv_state: jnp.ndarray | None = None) -> tuple[jnp.ndarray, jnp.ndarray]: + cache_len = self.kernel_size - 1 + + if conv_state is None: + x_padded = jnp.pad(x, ((0, 0), (cache_len, 0), (0, 0)), mode="constant", constant_values=0.0) + else: + x_padded = jnp.concatenate([jnp.transpose(conv_state, (0, 2, 1)), x], axis=1) + + output = self.conv(x_padded) + new_conv_state = jnp.transpose(x_padded[:, -cache_len:, :], (0, 2, 1)) + + return output, new_conv_state class Mamba2Mixer(nnx.Module): @@ -292,8 +345,11 @@ def __init__(self, cfg: Mamba2Config, layer_idx: int, *, rngs: nnx.Rngs): @jax.named_scope("mamba2_mixer") def __call__( - self, hidden_states: jnp.ndarray, initial_state: jnp.ndarray | None = None, return_final_state: bool = False - ) -> tuple[jnp.ndarray, jnp.ndarray | None]: + self, + hidden_states: jnp.ndarray, + conv_state: jnp.ndarray | None = None, + ssm_state: jnp.ndarray | None = None, + ) -> tuple[jnp.ndarray, jnp.ndarray, jnp.ndarray]: B_size, L, _ = hidden_states.shape # 1) Parallel projection @@ -311,18 +367,19 @@ def __call__( axis=-1, ) - # 2) Depthwise causal convolution - xBC = self.act(self.conv1d(xBC)) + # 2) Depthwise causal convolution with state caching + xBC, new_conv_state = self.conv1d(xBC, conv_state=conv_state) + xBC = self.act(xBC) x, B_t, C_t = jnp.split(xBC, [self.intermediate_size, self.intermediate_size + self.ssm_state_size], axis=-1) - # 3) SSD forward - init_state = initial_state[:, None, ...] if initial_state is not None else None + # 3) SSD forward with state caching + init_state = ssm_state[:, None, ...] if ssm_state is not None else None A = -jnp.exp(self.A_log[:].astype(jnp.float32)) B_exp = jnp.broadcast_to(jnp.expand_dims(B_t, 2), (B_size, L, self.num_heads, self.ssm_state_size)) C_exp = jnp.broadcast_to(jnp.expand_dims(C_t, 2), (B_size, L, self.num_heads, self.ssm_state_size)) - y, final_state = ssd_forward( + y, new_ssm_state = ssd_forward( x=x.reshape(B_size, L, -1, self.head_dim), dt=dt, A=A, @@ -334,7 +391,7 @@ def __call__( dt_min=self.dt_min, dt_max=self.dt_max, initial_states=init_state, - return_final_states=return_final_state, + return_final_states=True, ) y = y.reshape(B_size, L, -1) @@ -344,7 +401,7 @@ def __call__( y = jnp.concatenate([self.act(z0) * x0, y], axis=-1) # 5) Output projection - return self.out_proj(y), final_state + return self.out_proj(y), new_conv_state, new_ssm_state class Mamba2Block(nnx.Module): @@ -357,14 +414,17 @@ def __init__(self, cfg: Mamba2Config, layer_idx: int, *, rngs: nnx.Rngs): self.mixer = Mamba2Mixer(cfg, layer_idx=layer_idx, rngs=rngs) def __call__( - self, hidden_states: jnp.ndarray, initial_state: jnp.ndarray | None = None, return_final_state: bool = False - ) -> tuple[jnp.ndarray, jnp.ndarray | None]: + self, + hidden_states: jnp.ndarray, + conv_state: jnp.ndarray | None = None, + ssm_state: jnp.ndarray | None = None, + ) -> tuple[jnp.ndarray, jnp.ndarray, jnp.ndarray]: residual = hidden_states hs = self.norm(hidden_states.astype(jnp.float32)) if self.residual_in_fp32: residual = residual.astype(jnp.float32) - hs_out, last_state = self.mixer(hs, initial_state=initial_state, return_final_state=return_final_state) - return residual + hs_out, last_state + hs_out, new_conv_state, new_ssm_state = self.mixer(hs, conv_state=conv_state, ssm_state=ssm_state) + return residual + hs_out, new_conv_state, new_ssm_state class Mamba2Model(nnx.Module): @@ -381,40 +441,50 @@ def __call__( self, input_ids: jnp.ndarray | None = None, inputs_embeds: jnp.ndarray | None = None, - initial_states: list[jnp.ndarray] | None = None, + cache: Mamba2Cache | None = None, output_hidden_states: bool = False, - output_last_ssm_states: bool = False, - ) -> dict[str, jnp.ndarray | list[jnp.ndarray] | None]: + ) -> dict[str, jnp.ndarray | Mamba2Cache | list[jnp.ndarray] | None]: if (input_ids is None) == (inputs_embeds is None): raise ValueError("Specify exactly one of input_ids or inputs_embeds") hidden_states = self.embedder(input_ids) if inputs_embeds is None else inputs_embeds - if initial_states is None: - initial_states = [None] * self.cfg.num_hidden_layers - elif len(initial_states) != self.cfg.num_hidden_layers: - raise ValueError("initial_states length must equal num_hidden_layers") + # Extract per-layer states from cache or initialize + if cache is None: + conv_states = [None] * self.cfg.num_hidden_layers + ssm_states = [None] * self.cfg.num_hidden_layers + else: + if len(cache.conv_states) != self.cfg.num_hidden_layers: + raise ValueError("cache.conv_states length must equal num_hidden_layers") + if len(cache.ssm_states) != self.cfg.num_hidden_layers: + raise ValueError("cache.ssm_states length must equal num_hidden_layers") + conv_states = cache.conv_states + ssm_states = cache.ssm_states all_hidden_states = [] if output_hidden_states else None - all_last_states = [] if output_last_ssm_states else None + new_conv_states = [] + new_ssm_states = [] - for layer, init_state in zip(self.layers, initial_states): - hidden_states, last_state = layer( - hidden_states, initial_state=init_state, return_final_state=output_last_ssm_states + for layer, conv_state, ssm_state in zip(self.layers, conv_states, ssm_states): + hidden_states, new_conv_state, new_ssm_state = layer( + hidden_states, conv_state=conv_state, ssm_state=ssm_state ) + new_conv_states.append(new_conv_state) + new_ssm_states.append(new_ssm_state) if output_hidden_states: all_hidden_states.append(hidden_states) - if output_last_ssm_states: - all_last_states.append(last_state) hidden_states = self.final_norm(hidden_states) if output_hidden_states: all_hidden_states.append(hidden_states) + # Build updated cache + updated_cache = Mamba2Cache(ssm_states=new_ssm_states, conv_states=new_conv_states) + return { "last_hidden_state": hidden_states, + "cache": updated_cache, "hidden_states": all_hidden_states, - "last_ssm_states": all_last_states, } @@ -430,8 +500,13 @@ def __init__(self, cfg: Mamba2Config, *, rngs: nnx.Rngs): self.lm_head = None @jax.named_scope("mamba2_causal_lm") - def __call__(self, input_ids: jnp.ndarray, labels: jnp.ndarray | None = None) -> dict[str, jnp.ndarray | None]: - backbone_outputs = self.backbone(input_ids=input_ids) + def __call__( + self, + input_ids: jnp.ndarray, + labels: jnp.ndarray | None = None, + cache: Mamba2Cache | None = None, + ) -> dict[str, jnp.ndarray | Mamba2Cache | None]: + backbone_outputs = self.backbone(input_ids=input_ids, cache=cache) hidden_states = backbone_outputs["last_hidden_state"] if self.cfg.tie_word_embeddings: @@ -445,7 +520,7 @@ def __call__(self, input_ids: jnp.ndarray, labels: jnp.ndarray | None = None) -> shift_labels = labels[:, 1:].reshape(-1) loss = optax.softmax_cross_entropy_with_integer_labels(shift_logits, shift_labels).mean() - return {"logits": logits, "loss": loss} + return {"logits": logits, "loss": loss, "cache": backbone_outputs["cache"]} @classmethod def from_pretrained( @@ -513,6 +588,11 @@ def __call__(self, x: jnp.ndarray) -> jnp.ndarray: @jax.jit -def forward(model: Mamba2ForCausalLM, input_ids: jnp.ndarray, labels: jnp.ndarray | None = None): - """JIT-compiled forward pass for Mamba2ForCausalLM.""" - return model(input_ids, labels) +def forward( + model: Mamba2ForCausalLM, + input_ids: jnp.ndarray, + labels: jnp.ndarray | None = None, + cache: Mamba2Cache | None = None, +): + """JIT-compiled forward pass for Mamba2ForCausalLM with optional caching.""" + return model(input_ids, labels, cache) diff --git a/bonsai/models/mamba2/tests/run_model.py b/bonsai/models/mamba2/tests/run_model.py index a4242f2a..65d47209 100644 --- a/bonsai/models/mamba2/tests/run_model.py +++ b/bonsai/models/mamba2/tests/run_model.py @@ -19,7 +19,7 @@ from jax.sharding import PartitionSpec as P from transformers import AutoTokenizer -from bonsai.models.mamba2 import modeling +from bonsai.models.mamba2 import modeling, params @jax.jit @@ -73,8 +73,78 @@ def _to_host_1d(x: jnp.ndarray) -> jnp.ndarray: return jax.device_get(x) -def run_model(*, max_new_tokens: int = 32) -> None: - """Run the Mamba2 generation smoke test.""" +@jax.jit +def _cached_decode_step( + model: modeling.Mamba2ForCausalLM, + token: jnp.ndarray, + cache: modeling.Mamba2Cache, +) -> tuple[jnp.ndarray, modeling.Mamba2Cache]: + """Single decode step with cache. + + Args: + model: The Mamba2 model + token: Single token (batch, 1) + cache: Current cache state + + Returns: + next_token: Next predicted token (batch,) + new_cache: Updated cache + """ + out = model(token, labels=None, cache=cache) + logits = out["logits"][:, -1, :] + next_token = jnp.argmax(logits, axis=-1) + return next_token, out["cache"] + + +def _greedy_generate_cached( + model: modeling.Mamba2ForCausalLM, + prompt_ids_1d: jnp.ndarray, + *, + max_new_tokens: int, + cfg: modeling.Mamba2Config, +) -> jnp.ndarray: + """Greedy generation with SSM state caching (O(n) complexity). + + Args: + model: Mamba2ForCausalLM model + prompt_ids_1d: Prompt token IDs (seq_len,) + max_new_tokens: Number of tokens to generate + cfg: Model config for cache initialization + + Returns: + Generated tokens including prompt (1, prompt_len + max_new_tokens) + """ + prompt_len = int(prompt_ids_1d.shape[0]) + + # Prefill: Process prompt in one pass to initialize cache + prompt_2d = jnp.expand_dims(prompt_ids_1d, axis=0) # (1, prompt_len) + out = model(prompt_2d, labels=None, cache=None) + cache = out["cache"] + + # Get first generated token from prefill + logits = out["logits"][:, -1, :] + next_token = jnp.argmax(logits, axis=-1, keepdims=True) # (1, 1) + + generated = [prompt_2d, next_token] + + # Decode: Generate remaining tokens one at a time using cache + for _ in range(1, max_new_tokens): + next_token, cache = _cached_decode_step(model, next_token, cache) + next_token = jnp.expand_dims(next_token, axis=-1) # (1, 1) + generated.append(next_token) + + result = jnp.concatenate(generated, axis=1) + jax.block_until_ready(result) + return result + + +def run_model(*, max_new_tokens: int = 32, use_cache: bool = True) -> None: + """Run the Mamba2 generation smoke test. + + Args: + max_new_tokens: Number of tokens to generate + use_cache: If True, use cached generation (O(n)). If False, use non-cached (O(n²)) + """ query = [ "Why is the sky blue instead of any other color like purple?", "What is the capital city of England?", @@ -93,40 +163,56 @@ def run_model(*, max_new_tokens: int = 32) -> None: model = modeling.Mamba2ForCausalLM.from_pretrained("state-spaces/mamba2-130m", cfg=cfg) tokenizer = AutoTokenizer.from_pretrained("EleutherAI/gpt-neox-20b") - # Determine padding id robustly. - pad_id = tokenizer.pad_token_id - if pad_id is None: - pad_id = tokenizer.eos_token_id - if pad_id is None: - pad_id = cfg.pad_token_id - - # Share a single static buffer length across all prompts to reuse compilation. prompt_ids_list = [jnp.asarray(tokenizer.encode(q), dtype=jnp.int32) for q in query] - max_prompt_len = max(int(x.shape[0]) for x in prompt_ids_list) - buffer_len = max_prompt_len + max_new_tokens - - for q, prompt_ids in zip(query, prompt_ids_list): - tokens_2d = _greedy_generate( - model, - prompt_ids, - max_new_tokens=max_new_tokens, - pad_id=int(pad_id), - buffer_len=buffer_len, - ) - - host_ids = _to_host_1d(tokens_2d.at[0].get()) - - generated_ids_only = host_ids[len(prompt_ids) :] - - text = tokenizer.decode(generated_ids_only.tolist(), skip_special_tokens=True) - print(f"User:\n {q}") - print(f"Answer:\n {text.strip()}\n\n") + if use_cache: + print("Using cached generation (O(n) complexity)\n") + for q, prompt_ids in zip(query, prompt_ids_list): + tokens_2d = _greedy_generate_cached( + model, + prompt_ids, + max_new_tokens=max_new_tokens, + cfg=cfg, + ) + + host_ids = _to_host_1d(tokens_2d.at[0].get()) + generated_ids_only = host_ids[len(prompt_ids) :] + text = tokenizer.decode(generated_ids_only.tolist(), skip_special_tokens=True) + + print(f"User:\n {q}") + print(f"Answer:\n {text.strip()}\n") + else: + print("Using non-cached generation (O(n^2) complexity)\n") + # Determine padding id robustly. + pad_id = tokenizer.pad_token_id + if pad_id is None: + pad_id = tokenizer.eos_token_id + if pad_id is None: + pad_id = cfg.pad_token_id + + # Share a single static buffer length across all prompts to reuse compilation. + max_prompt_len = max(int(x.shape[0]) for x in prompt_ids_list) + buffer_len = max_prompt_len + max_new_tokens + + for q, prompt_ids in zip(query, prompt_ids_list): + tokens_2d = _greedy_generate( + model, + prompt_ids, + max_new_tokens=max_new_tokens, + pad_id=int(pad_id), + buffer_len=buffer_len, + ) + + host_ids = _to_host_1d(tokens_2d.at[0].get()) + generated_ids_only = host_ids[len(prompt_ids) :] + text = tokenizer.decode(generated_ids_only.tolist(), skip_special_tokens=True) + + print(f"User:\n {q}") + print(f"Answer:\n {text.strip()}\n") def run_forecaster() -> None: """Run a tiny Mamba2Forecaster smoke test (shape-only).""" - from bonsai.models.mamba2 import params model = params.create_random_forecaster( input_dim=10, diff --git a/bonsai/models/mamba2/tests/test_outputs_mamba_2.py b/bonsai/models/mamba2/tests/test_outputs_mamba_2.py index 5597e7bd..2334ff05 100644 --- a/bonsai/models/mamba2/tests/test_outputs_mamba_2.py +++ b/bonsai/models/mamba2/tests/test_outputs_mamba_2.py @@ -166,7 +166,8 @@ def test_output_shape(self): self.assertEqual(outputs["last_hidden_state"].shape, (batch_size, seq_len, self.cfg.hidden_size)) self.assertIsNone(outputs["hidden_states"]) - self.assertIsNone(outputs["last_ssm_states"]) + self.assertIsNotNone(outputs["cache"]) + self.assertIsInstance(outputs["cache"], modeling.Mamba2Cache) def test_output_hidden_states(self): """Test output_hidden_states flag.""" @@ -216,6 +217,7 @@ def test_output_shape(self): self.assertEqual(outputs["logits"].shape, (batch_size, seq_len, self.cfg.vocab_size)) self.assertIsNone(outputs["loss"]) + self.assertIsNotNone(outputs["cache"]) def test_loss_computation(self): """Test loss computation with labels.""" @@ -409,5 +411,102 @@ def test_logits_parity(self): np.testing.assert_allclose(bonsai_logits, golden_logits, rtol=rtol, atol=atol) +class TestMamba2Cache(absltest.TestCase): + """Tests for Mamba2Cache state caching.""" + + def setUp(self): + super().setUp() + self.cfg = modeling.Mamba2Config.tiny() + self.model = modeling.Mamba2ForCausalLM(self.cfg, rngs=nnx.Rngs(42)) + + def test_cache_shapes(self): + """Test cache state shapes are correct.""" + batch_size, seq_len = 2, 32 + input_ids = jnp.ones((batch_size, seq_len), dtype=jnp.int32) + outputs = self.model(input_ids=input_ids) + cache = outputs["cache"] + + # Check cache structure + self.assertIsInstance(cache, modeling.Mamba2Cache) + self.assertLen(cache.ssm_states, self.cfg.num_hidden_layers) + self.assertLen(cache.conv_states, self.cfg.num_hidden_layers) + + # Check SSM state shapes + for ssm_state in cache.ssm_states: + expected_shape = (batch_size, self.cfg.num_heads, self.cfg.head_dim, self.cfg.state_size) + self.assertEqual(ssm_state.shape, expected_shape) + + # Check conv state shapes + conv_dim = self.cfg.intermediate_size + 2 * self.cfg.state_size + cache_len = self.cfg.conv_kernel - 1 + for conv_state in cache.conv_states: + expected_shape = (batch_size, conv_dim, cache_len) + self.assertEqual(conv_state.shape, expected_shape) + + def test_cached_matches_full(self): + """Test that cached generation produces same results as full forward pass. + + Process [a,b,c] in full sequence vs [a,b] then [c] with cache. + Logits for 'c' position must match. + """ + batch_size = 1 + # Create sequence [1, 2, 3] + full_seq = jnp.array([[1, 2, 3]], dtype=jnp.int32) + prefix_seq = jnp.array([[1, 2]], dtype=jnp.int32) + next_token = jnp.array([[3]], dtype=jnp.int32) + + # Full forward pass + full_outputs = self.model(input_ids=full_seq) + full_logits = full_outputs["logits"][:, -1, :] # logits at position of token 3 + + # Cached forward pass + prefix_outputs = self.model(input_ids=prefix_seq) + cache = prefix_outputs["cache"] + next_outputs = self.model(input_ids=next_token, cache=cache) + cached_logits = next_outputs["logits"][:, -1, :] # logits at position of token 3 + + # Logits should match + np.testing.assert_allclose(np.array(full_logits), np.array(cached_logits), rtol=1e-5, atol=1e-6) + + def test_create_empty_cache(self): + """Test creating empty cache with correct shapes.""" + batch_size = 4 + cache = modeling.create_empty_cache(self.cfg, batch_size) + + self.assertIsInstance(cache, modeling.Mamba2Cache) + self.assertLen(cache.ssm_states, self.cfg.num_hidden_layers) + self.assertLen(cache.conv_states, self.cfg.num_hidden_layers) + + # Check all states are zeros with correct shapes + for ssm_state in cache.ssm_states: + expected_shape = (batch_size, self.cfg.num_heads, self.cfg.head_dim, self.cfg.state_size) + self.assertEqual(ssm_state.shape, expected_shape) + self.assertTrue(jnp.all(ssm_state == 0)) + + conv_dim = self.cfg.intermediate_size + 2 * self.cfg.state_size + cache_len = self.cfg.conv_kernel - 1 + for conv_state in cache.conv_states: + expected_shape = (batch_size, conv_dim, cache_len) + self.assertEqual(conv_state.shape, expected_shape) + self.assertTrue(jnp.all(conv_state == 0)) + + def test_cache_updates_on_forward(self): + """Test that cache is updated after each forward pass.""" + input_ids = jnp.array([[1, 2]], dtype=jnp.int32) + + # First forward pass + outputs1 = self.model(input_ids=input_ids) + cache1 = outputs1["cache"] + + # Second forward pass with cache + next_token = jnp.array([[3]], dtype=jnp.int32) + outputs2 = self.model(input_ids=next_token, cache=cache1) + cache2 = outputs2["cache"] + + # Caches should be different (states updated) + for s1, s2 in zip(cache1.ssm_states, cache2.ssm_states): + self.assertFalse(jnp.allclose(s1, s2)) + + if __name__ == "__main__": absltest.main() From 077f43b5ce504f3412e00d5aec3e83ab4b7aee69 Mon Sep 17 00:00:00 2001 From: Jen Ha <25069493+jenriver@users.noreply.github.com> Date: Thu, 22 Jan 2026 13:49:52 -0800 Subject: [PATCH 18/18] Temporarily disable Gemini agents until specific review guidelines are configured. (#140) --- .gemini/config.yaml | 7 +++++++ 1 file changed, 7 insertions(+) create mode 100644 .gemini/config.yaml diff --git a/.gemini/config.yaml b/.gemini/config.yaml new file mode 100644 index 00000000..9142ef92 --- /dev/null +++ b/.gemini/config.yaml @@ -0,0 +1,7 @@ +version: 1 +enabled_features: + auto_review: + enabled: false + drafts: false + review_pull_request: false + post_review_summary: false