""" rec7_model.py — the rec7 substitute encoder for MiniMax-Music3. Audio -> Flow-VAE latents -> per-frame hidden states (4096-d @ 25 Hz) that the M3 renderer accepts as if the language model had produced them. From those states the 8 code streams and the full renderer condition follow deterministically (semantic head argmax + the released depth decoder's greedy chain), which is how "covers" work: real recording in, M3 renders it back, timbre and performance preserved. Everything here is inference-only and self-contained apart from: - the MiniMax-Music3 checkpoint (Flow-VAE, depth decoder, LM tables) - dav_loader.py (thin loader for the Flow-VAE; ships alongside) """ from __future__ import annotations import json from pathlib import Path import numpy as np import torch import torch.nn as nn import torch.nn.functional as F # ── geometry (do not change: it matches the renderer's windowing) ── LATENT_CHANNELS = 128 FRAMES = 128 # encoder window, in 25 Hz frames LATENT_WINDOW_MAX = 448 RATIO_NUM, RATIO_DEN = 441, 128 CHUNK_FRAMES, CHUNK_HOP = 200, 100 HOP_LATENTS = 345 OWNED_FROM = 25 SEM_VOCAB, AC_VOCAB, N_AC = 16384, 1024, 7 SEM_OFFSET = 151675 # c0 lives here in the LM vocabulary H_DIM = 4096 def n_dit_windows(n_frames): return max(1, (n_frames - 1) // CHUNK_HOP) def frame_latent_starts(n_frames): """Frame t -> first latent index. Piecewise, NOT t*441/128: the renderer denoises 200-frame windows on a 100-frame hop stitched at an integer 345-latent hop, and non-first windows take ownership 25 frames after their nominal start.""" t = np.arange(n_frames + 1, dtype=np.int64) k = np.clip((t - OWNED_FROM) // CHUNK_HOP, 0, n_dit_windows(n_frames) - 1) tau = t - k * CHUNK_HOP Fw = np.minimum(CHUNK_FRAMES, n_frames - k * CHUNK_HOP) L = Fw * RATIO_NUM // RATIO_DEN return k * HOP_LATENTS + (tau * L + Fw - 1) // Fw def pool_matrix(bounds): pool = np.zeros((FRAMES, LATENT_WINDOW_MAX), dtype=np.float32) for j in range(FRAMES): a, b = int(bounds[j]), int(bounds[j + 1]) if b > a: pool[j, a:b] = 1.0 / (b - a) return pool # ── architecture ── class ResBlock(nn.Module): def __init__(self, d, dilation): super().__init__() self.norm = nn.GroupNorm(1, d) self.conv1 = nn.Conv1d(d, d, 3, padding=dilation, dilation=dilation) self.conv2 = nn.Conv1d(d, d, 1) def forward(self, x): h = self.conv1(F.gelu(self.norm(x))) return x + self.conv2(F.gelu(h)) class DepthDecoder(nn.Module): """Internal auxiliary head used during training; kept so the checkpoint loads. Not used at inference (the released M3 depth decoder produces the acoustic codes from the states).""" def __init__(self, d_ctx=1088, d=512, layers=2, heads=8, ff_mult=4, dropout=0.1): super().__init__() self.proj = nn.Linear(d_ctx, d) self.sem_emb = nn.Embedding(SEM_VOCAB, d) self.ac_emb = nn.Embedding((N_AC - 1) * AC_VOCAB, d) self.pos = nn.Parameter(torch.zeros(1, 8, d)) layer = nn.TransformerEncoderLayer( d, heads, d * ff_mult, dropout=dropout, activation="gelu", batch_first=True, norm_first=True) self.tr = nn.TransformerEncoder(layer, layers) self.heads = nn.ModuleList(nn.Linear(d, AC_VOCAB) for _ in range(N_AC)) mask = torch.triu(torch.full((8, 8), float("-inf")), diagonal=1) self.register_buffer("causal", mask, persistent=False) class V4Encoder(nn.Module): def __init__(self, d_model=1088, n_layers=8, n_heads=17, ff_mult=4, dropout=0.1, depth_d=512, depth_layers=2, depth_heads=8): super().__init__() self.conv_in = nn.Conv1d(LATENT_CHANNELS, d_model, 7, padding=3) self.blocks = nn.ModuleList(ResBlock(d_model, d) for d in (1, 3, 9)) self.pos = nn.Parameter(torch.zeros(1, FRAMES, d_model)) layer = nn.TransformerEncoderLayer( d_model, n_heads, d_model * ff_mult, dropout=dropout, activation="gelu", batch_first=True, norm_first=True) self.transformer = nn.TransformerEncoder(layer, n_layers) self.norm_out = nn.LayerNorm(d_model) self.sem_head = nn.Linear(d_model, SEM_VOCAB) self.depth = DepthDecoder(d_model, depth_d, depth_layers, depth_heads, ff_mult, dropout) def features(self, latents, pool): x = self.conv_in(latents.transpose(1, 2)) for b in self.blocks: x = b(x) x = torch.bmm(pool, x.transpose(1, 2)) + self.pos return self.norm_out(self.transformer(x)) # [B, 128, 1088] def forward(self, latents, pool): feats = self.features(latents, pool) return feats, self.sem_head(feats) class HHead(nn.Module): def __init__(self, d_in=1088): super().__init__() self.net = nn.Sequential(nn.Linear(d_in, 2048), nn.GELU(), nn.Linear(2048, H_DIM)) def forward(self, x): return self.net(x) # ── loading ── def load_rec7(weights_dir, device="cuda"): weights_dir = Path(weights_dir) trunk = V4Encoder().to(device).eval() trunk.load_state_dict(torch.load(weights_dir / "trunk.pt", map_location="cpu", weights_only=False)["model"]) head = HHead(trunk.pos.shape[-1]).to(device).eval() head.load_state_dict(torch.load(weights_dir / "head.pt", map_location="cpu", weights_only=False)["model"]) for p in list(trunk.parameters()) + list(head.parameters()): p.requires_grad_(False) return trunk, head def load_dav(m3_dir, device="cuda"): """Flow-VAE from the M3 checkpoint, via the small loader shipped here.""" from dav_loader import get_dav_class m3_dir = Path(m3_dir) dav = get_dav_class(m3_dir.parent).from_original_dav( str(m3_dir / "dav.pth")).to(device).eval() for p in dav.parameters(): p.requires_grad_(False) return dav def lm_semantic_tables(m3_dir, device="cuda"): """The 16k-row slices of the LM's output head and input embedding that cover the semantic codes. Read directly from safetensors, so the 8B never has to be loaded.""" from safetensors import safe_open m3_dir = Path(m3_dir) cands = [m3_dir / "language_model", m3_dir / "qwen_7B" / "qwen_7B"] root = next((c for c in cands if (c / "model.safetensors.index.json") .exists() or (c / "model.safetensors").exists()), None) if root is None: raise FileNotFoundError("language model safetensors not found " f"under {m3_dir}") idx = root / "model.safetensors.index.json" if idx.exists(): wm = json.loads(idx.read_text())["weight_map"] def where(suffix): key = next(k for k in wm if k.endswith(suffix)) return root / wm[key], key else: with safe_open(str(root / "model.safetensors"), framework="pt") as f: keys = list(f.keys()) def where(suffix): return root / "model.safetensors", \ next(k for k in keys if k.endswith(suffix)) def sl(suffix): fpath, key = where(suffix) with safe_open(str(fpath), framework="pt") as f: return f.get_slice(key)[SEM_OFFSET: SEM_OFFSET + SEM_VOCAB] \ .to(device, torch.bfloat16) return sl("lm_head.weight"), sl("embed_tokens.weight") # ── inference ── @torch.no_grad() def encode_audio(dav, wav, device="cuda", chunk_s=12.0, overlap_s=2.0, sr=44100, hop=512): """wav [2, N] float32 @ 44.1 kHz -> latents [128, L] @ 86.13 Hz.""" C, O = int(chunk_s * sr), (int(overlap_s * sr) // hop) * hop outs, pos, N = [], 0, wav.shape[-1] while pos < N: end = min(pos + C, N) st = max(0, pos - O) z = dav.encode(wav[..., st:end].unsqueeze(0).to(device))[0] outs.append(z[:, (pos - st) // hop:].cpu()) pos = end return torch.cat(outs, -1) @torch.no_grad() def read_states(trunk, head, z, device="cuda", max_frames=None): """latents [128, L] -> states [T, 4096] @ 25 Hz, stitched from overlapping 128-frame windows (hop 72).""" z = z.to(device).float() T_tot = int(z.shape[1] / 3.45) - 2 T = T_tot - 1 if max_frames is None else min(max_frames, T_tot - 1) if T < FRAMES + 2: raise ValueError("audio too short (need > ~5.5 s)") st = frame_latent_starts(T_tot) hop = max(1, (FRAMES * 9) // 16) acc = torch.zeros(T, H_DIM, device=device) cnt = torch.zeros(T, 1, device=device) offs = list(range(0, max(1, T - FRAMES + 1), hop)) if offs[-1] != T - FRAMES: offs.append(max(0, T - FRAMES)) for o in offs: b = st[o: o + FRAMES + 1] - st[o] n = int(b[-1]) if n > LATENT_WINDOW_MAX or st[o] + n > z.shape[1]: continue lat = torch.zeros(1, LATENT_WINDOW_MAX, LATENT_CHANNELS) lat[0, :n] = z[:, int(st[o]): int(st[o]) + n].T pl = torch.from_numpy(pool_matrix(b)).unsqueeze(0) with torch.autocast("cuda", dtype=torch.bfloat16): feats, _ = trunk(lat.to(device), pl.to(device)) hw = head(feats[0].float()) e = min(o + FRAMES, T) acc[o: e] += hw[: e - o] cnt[o: e] += 1 return acc / cnt.clamp_min(1) @torch.no_grad() def states_to_streams(h, depth, W_head, W_embed, device="cuda", chunk=512): """states [T, 4096] -> (codes [T, 8] long, condition [T, 32768]) using the released M3 depth decoder: c0 by semantic-head argmax, c1..c7 by the greedy chain, and the 8 hidden streams the renderer's condition encoder expects, concatenated per frame.""" codes_out, cond_out = [], [] for s in range(0, h.shape[0], chunk): hb = h[s: s + chunk].to(device, torch.bfloat16) c0 = (hb @ W_head.T).float().argmax(-1) seq = [depth.projection(hb).unsqueeze(1), depth.projection(W_embed[c0]).unsqueeze(1)] hiddens, cols = [hb], [c0] for k in range(1, 8): hid = depth(torch.cat(seq, dim=1))[:, -1] hiddens.append(hid) lg = depth.audio_heads[k - 1](hid).float() idx = lg.argmax(-1).clamp(0, AC_VOCAB - 1) cols.append(idx) if k < 7: emb = depth.audio_embeddings.weight[(k - 1) * AC_VOCAB + idx] seq.append(depth.projection(emb.to(hb.dtype)).unsqueeze(1)) codes_out.append(torch.stack(cols, dim=-1)) cond_out.append(torch.cat(hiddens, dim=-1)) return torch.cat(codes_out).cpu(), torch.cat(cond_out)