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) diff --git a/bonsai/models/mimo_audio/__init__.py b/bonsai/models/mimo_audio/__init__.py new file mode 100644 index 00000000..5f102705 --- /dev/null +++ b/bonsai/models/mimo_audio/__init__.py @@ -0,0 +1,19 @@ +from bonsai.models.mimo_audio.modeling import ( + MiMoAudioConfig, + MiMoAudioArguments, + FlaxMiMoAudioForCausalLM, +) +from bonsai.models.mimo_audio.mimo_audio_tokenizer import ( + FlaxMiMoAudioTokenizer, + MiMoAudioTokenizerConfig, + MelSpectrogram, +) + +__all__ = [ + "MiMoAudioConfig", + "MiMoAudioArguments", + "FlaxMiMoAudioForCausalLM", + "FlaxMiMoAudioTokenizer", + "MiMoAudioTokenizerConfig", + "MelSpectrogram", +] 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..75ec1811 --- /dev/null +++ b/bonsai/models/mimo_audio/mimo_audio_configuration.py @@ -0,0 +1,110 @@ +from dataclasses import dataclass +from typing import TYPE_CHECKING +from bonsai.models.qwen3.modeling import ShardingCfg + +if TYPE_CHECKING: + from bonsai.models.qwen2.modeling import ModelConfig as Qwen2Config + + +@dataclass +class MiMoAudioConfig: + 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_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_layers: int = 6 + input_local_dim: int = 1024 + input_full_attention: bool = True + + 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 + + return Qwen2Config( + num_layers=6, + vocab_size=151680, + emb_dim=1024, + mlp_dim=4096, + num_heads=64, + head_dim=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: + 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: + do_sample: bool = True + temperature: float = 1.0 + top_k: int = 50 + top_p: float = 0.95 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..6bc88f2a --- /dev/null +++ b/bonsai/models/mimo_audio/mimo_audio_tokenizer.py @@ -0,0 +1,844 @@ +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: + 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): + 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=1, + 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=1, + 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..111d9b9b --- /dev/null +++ b/bonsai/models/mimo_audio/mimo_audio_tokenizer_configuration.py @@ -0,0 +1,210 @@ +"""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: + 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..0af1076f --- /dev/null +++ b/bonsai/models/mimo_audio/mimo_audio_tokenizer_params.py @@ -0,0 +1,372 @@ +"""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 diff --git a/bonsai/models/mimo_audio/modeling.py b/bonsai/models/mimo_audio/modeling.py new file mode 100644 index 00000000..877cbd27 --- /dev/null +++ b/bonsai/models/mimo_audio/modeling.py @@ -0,0 +1,295 @@ +from typing import Optional, Tuple, List +import jax +import jax.numpy as jnp +from flax import nnx +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 + + +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]: + 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) + + 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) + + 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, + 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 + + +@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]: + 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]: + 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: + return model.local_forward(local_embeds, key, local_sampler=None) diff --git a/bonsai/models/mimo_audio/params.py b/bonsai/models/mimo_audio/params.py new file mode 100644 index 00000000..7cb2392b --- /dev/null +++ b/bonsai/models/mimo_audio/params.py @@ -0,0 +1,179 @@ +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 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..9d50970a --- /dev/null +++ b/bonsai/models/mimo_audio/test/run_model.py @@ -0,0 +1,240 @@ +#!/usr/bin/env python3 + +import os +import json +import jax +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): + 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(): + 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) + + 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() 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..eb852cca --- /dev/null +++ b/bonsai/models/mimo_audio/test/test_outputs_mimo_audio.py @@ -0,0 +1,394 @@ +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 + +# 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 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/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..e35959ad --- /dev/null +++ b/bonsai/models/mimo_audio/test/test_outputs_mimo_audio_tokenizer.py @@ -0,0 +1,186 @@ +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() 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 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..f6449632 --- /dev/null +++ b/bonsai/models/qwen2/modeling.py @@ -0,0 +1,275 @@ +# 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 get_abstract_mesh +from jaxtyping import Array + +from bonsai.models.qwen3.modeling import ( + Cache, + LayerCache, + MLP, + RMSNorm, + ShardingCfg, + _generate_pos_embeddings, + apply_rope, + compute_positions_from_segment_ids, + count_left_pads, + count_right_pads, + 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 + + 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 + ) + 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.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.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: + 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] + + 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) + + 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) + + 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 + + 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 + + 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 = kv_segment_ids[:, None, :] == segment_ids[:, :, None] + + if self.cfg.use_causal_mask: + causal_mask = k_pos[:, None, :] <= q_pos[:, :, None] + final_mask = causal_mask & segment_mask + else: + final_mask = segment_mask + + attn_mask = final_mask[:, :, :, None, None] + attn_logits = jnp.where(attn_mask, attn_logits, _K_MASK) + + 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)) + + 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) + 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)) + + if not get_abstract_mesh().empty: + 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 diff --git a/bonsai/models/qwen2/params.py b/bonsai/models/qwen2/params.py new file mode 100644 index 00000000..f474b187 --- /dev/null +++ b/bonsai/models/qwen2/params.py @@ -0,0 +1,139 @@ +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]]: + return { + r"model\.embed_tokens\.weight": ("embedder.embedding", TRANSFORM_NONE), + 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), + 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), + 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), + 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, + ), + 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 diff --git a/bonsai/models/qwen2/tests/__init__.py b/bonsai/models/qwen2/tests/__init__.py new file mode 100644 index 00000000..1337256a --- /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. diff --git a/bonsai/models/qwen2/tests/run_model.py b/bonsai/models/qwen2/tests/run_model.py new file mode 100644 index 00000000..198bd96a --- /dev/null +++ b/bonsai/models/qwen2/tests/run_model.py @@ -0,0 +1,98 @@ +import jax +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 +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) + 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_name = "Qwen/Qwen2-7B" + model_ckpt_path = snapshot_download(model_name) + + 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_name, + trust_remote_code=True, + ) + + print() + tokens = tokenize(tokenizer, query, batch_shd) + batch_size, token_len = tokens.shape + + generate_steps = 1024 + model = params.create_model_from_safe_tensors(model_ckpt_path, config, mesh) + + 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): + key, subkey = jax.random.split(key) + next_tokens = jit_sampler(logits, key=subkey) + + current_token_id = int(next_tokens.squeeze(-1)[0]) + + 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"] 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..25293451 --- /dev/null +++ b/bonsai/models/qwen2/tests/test_outputs_qwen2.py @@ -0,0 +1,365 @@ +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()