"""Encoder-decoder Transformer: log-mel windows -> chart token sequences.""" import math import torch import torch.nn as nn from .vocab import N_MELS, VOCAB, WINDOW def sinusoidal(length, dim): pos = torch.arange(length, dtype=torch.float32)[:, None] i = torch.arange(dim // 2, dtype=torch.float32)[None] angle = pos / torch.pow(10000.0, 2 * i / dim) emb = torch.zeros(length, dim) emb[:, 0::2] = torch.sin(angle) emb[:, 1::2] = torch.cos(angle) return emb class BeatHiRes(nn.Module): """Frame-resolution beat/downbeat head: encoder states (T/4) are upsampled 4x so supervision and peaks live on the 11.6 ms frame grid instead of 46 ms encoder bins (4x finer timing per peak).""" def __init__(self, d): super().__init__() self.up = nn.ConvTranspose1d(d, d // 2, 4, stride=4) self.out = nn.Conv1d(d // 2, 2, 3, padding=1) def forward(self, mem): # (B, L, d) -> (B, 4L, 2) h = nn.functional.gelu(self.up(mem.transpose(1, 2))) return self.out(h).transpose(1, 2) def _parse_pattern(spec, n): """Layer-sharing schedule. spec is either an int (number of UNIQUE physical layers, cycled to length n) or a comma string like "E0,E1,E0,E1" (labels; each distinct label is one physical layer). Returns a list of n physical-layer indices, e.g. [0,1,0,1] for 2-unique/4-physical-depth.""" if spec is None: return list(range(n)) # no sharing: every layer unique if isinstance(spec, int): return [k % spec for k in range(n)] labels = [s.strip() for s in str(spec).split(",") if s.strip()] assert len(labels) == n, f"pattern {spec} has {len(labels)} labels, need {n}" order = {} out = [] for lb in labels: if lb not in order: order[lb] = len(order) out.append(order[lb]) return out class Adapter(nn.Module): """Rank-r residual adapter applied after a (possibly reused) transformer block: x + gate * up(gelu(down(x))). gate init 0 so the whole stack starts identical to hard weight sharing; the adapter only differentiates a reused physical layer as training moves the gate off zero.""" def __init__(self, d, rank): super().__init__() self.down = nn.Linear(d, rank, bias=False) self.up = nn.Linear(rank, d, bias=False) nn.init.normal_(self.down.weight, std=0.02) nn.init.zeros_(self.up.weight) self.gate = nn.Parameter(torch.zeros(1)) def forward(self, x): return x + self.gate * self.up(nn.functional.gelu(self.down(x))) class SharedEncoder(nn.Module): """Encoder whose physical layers are reused according to a sharing pattern. Each POSITION in the depth keeps its own (a) rank-r adapter and (b) final LayerNorm — so ΔW starts at 0 (hard sharing) but every position can drift. With pattern=None this is exactly n unique layers + a norm, matching the stock nn.TransformerEncoder param count (adapters/pos-norms are opt-in). v1.7 depth-diversity options (all default-off, init == v1.6 behaviour): depth_emb: per-depth learned d-vector ADDED to the block input at each reuse (params = depth x d, zero-init so init == hard sharing). adapter_rank_ffn > 0: SPLIT adapters — instead of one post-block adapter, each depth owns a rank-`adapter_rank` adapter on the attention sublayer output and a rank-`adapter_rank_ffn` adapter on the FFN sublayer output (patterning lives in the FFN). Gate init 0 keeps init == hard sharing. """ def __init__(self, d, nhead, ffn, dropout, pattern, adapter_rank=0, unique_layernorm=False, adapter_rank_ffn=0, depth_emb=False): super().__init__() # `pattern` is the resolved list of physical-layer indices per depth n_phys = max(pattern) + 1 self.layers = nn.ModuleList([ nn.TransformerEncoderLayer(d, nhead, ffn, dropout, activation="gelu", batch_first=True, norm_first=True) for _ in range(n_phys)]) self.plan = pattern self.split = adapter_rank_ffn > 0 if self.split: # v1.7: per-sublayer adapters (attention rank-a, FFN rank-f) self.attn_adapters = nn.ModuleList([ Adapter(d, adapter_rank) if adapter_rank else nn.Identity() for _ in pattern]) self.ffn_adapters = nn.ModuleList([ Adapter(d, adapter_rank_ffn) for _ in pattern]) else: # v1.6: one residual adapter after the whole block self.adapters = nn.ModuleList([ Adapter(d, adapter_rank) if adapter_rank else nn.Identity() for _ in pattern]) self.depth_emb = (nn.Parameter(torch.zeros(len(pattern), d)) if depth_emb else None) self.pos_norms = nn.ModuleList([ nn.LayerNorm(d) if unique_layernorm else nn.Identity() for _ in pattern]) self.norm = nn.LayerNorm(d) def forward(self, x, src_key_padding_mask=None): for k, phys in enumerate(self.plan): if self.depth_emb is not None: x = x + self.depth_emb[k] layer = self.layers[phys] if self.split: # norm_first decomposition of nn.TransformerEncoderLayer with # per-depth adapters applied to each sublayer OUTPUT (Houlsby # placement); at gate=0 this is bit-equal to the stock layer x = x + self.attn_adapters[k]( layer._sa_block(layer.norm1(x), None, src_key_padding_mask)) x = x + self.ffn_adapters[k](layer._ff_block(layer.norm2(x))) else: x = layer(x, src_key_padding_mask=src_key_padding_mask) x = self.adapters[k](x) x = self.pos_norms[k](x) return self.norm(x) class SharedDecoder(nn.Module): """Decoder counterpart of SharedEncoder (see that docstring). Split mode puts rank-`adapter_rank` adapters on BOTH attention sublayers (self and cross) and rank-`adapter_rank_ffn` on the FFN sublayer.""" def __init__(self, d, nhead, ffn, dropout, pattern, adapter_rank=0, unique_layernorm=False, adapter_rank_ffn=0, depth_emb=False): super().__init__() n_phys = max(pattern) + 1 self.layers = nn.ModuleList([ nn.TransformerDecoderLayer(d, nhead, ffn, dropout, activation="gelu", batch_first=True, norm_first=True) for _ in range(n_phys)]) self.plan = pattern self.split = adapter_rank_ffn > 0 if self.split: self.sa_adapters = nn.ModuleList([ Adapter(d, adapter_rank) if adapter_rank else nn.Identity() for _ in pattern]) self.ca_adapters = nn.ModuleList([ Adapter(d, adapter_rank) if adapter_rank else nn.Identity() for _ in pattern]) self.ffn_adapters = nn.ModuleList([ Adapter(d, adapter_rank_ffn) for _ in pattern]) else: self.adapters = nn.ModuleList([ Adapter(d, adapter_rank) if adapter_rank else nn.Identity() for _ in pattern]) self.depth_emb = (nn.Parameter(torch.zeros(len(pattern), d)) if depth_emb else None) self.pos_norms = nn.ModuleList([ nn.LayerNorm(d) if unique_layernorm else nn.Identity() for _ in pattern]) self.norm = nn.LayerNorm(d) def forward(self, tgt, memory, tgt_mask=None, tgt_key_padding_mask=None, tgt_is_causal=False): x = tgt for k, phys in enumerate(self.plan): if self.depth_emb is not None: x = x + self.depth_emb[k] layer = self.layers[phys] if self.split: x = x + self.sa_adapters[k](layer._sa_block( layer.norm1(x), tgt_mask, tgt_key_padding_mask, tgt_is_causal)) x = x + self.ca_adapters[k](layer._mha_block( layer.norm2(x), memory, None, None)) x = x + self.ffn_adapters[k](layer._ff_block(layer.norm3(x))) else: x = layer(x, memory, tgt_mask=tgt_mask, tgt_key_padding_mask=tgt_key_padding_mask, tgt_is_causal=tgt_is_causal) x = self.adapters[k](x) x = self.pos_norms[k](x) return self.norm(x) class ChartModel(nn.Module): def __init__( self, d_model=512, nhead=8, enc_layers=6, dec_layers=6, ffn=2048, dropout=0.1, vocab_size=None, max_tgt=None, aux=False, global_ctx=False, func_time=False, ptr=False, beat=False, in_ch=None, emb_factor=None, clean_phase=False, enc_share=None, dec_share=None, adapter_rank=0, unique_layernorm=False, adapter_rank_ffn=0, depth_emb=False, unshare_last_dec=False, ): super().__init__() from .vocab import MAX_TGT vocab_size = vocab_size or VOCAB.size max_tgt = max_tgt or MAX_TGT self.d_model = d_model # clean_phase (v1.6): the encoder never sees the 2 grid-phase channels # (in_ch stays N_MELS), and phase is injected into the DECODER memory # only, via a small projection. This removes the condition leak that # forced beat supervision to be masked on slot windows in v1.5. self.clean_phase = clean_phase self.in_ch = in_ch or N_MELS # legacy slot mode appends 2 grid-phase channels front_ch = N_MELS if clean_phase else self.in_ch # conv frontend: (B, front_ch, T) -> (B, T/4, d) self.frontend = nn.Sequential( nn.Conv1d(front_ch, d_model, 3, stride=2, padding=1), nn.GELU(), nn.Conv1d(d_model, d_model, 3, stride=2, padding=1), nn.GELU(), ) enc_len = WINDOW // 4 self.register_buffer("enc_pos", sinusoidal(enc_len, d_model), persistent=False) self.register_buffer("dec_pos", sinusoidal(max_tgt, d_model), persistent=False) # phase injection (clean_phase only): 2 grid-phase channels at frame # rate -> encoder rate (T/4) -> additive conditioning on the memory the # DECODER reads. The beat head reads the pre-injection (clean) memory. # phase_proj input = [measure_phase, beat_phase, valid_flag]. valid=0 on # time-mode windows (their raw phase is -1), giving the decoder a clean # zero-phase conditioning there; slot windows pass valid=1 + real phase. if clean_phase: self.phase_proj = nn.Sequential( nn.Conv1d(3, d_model, 3, stride=2, padding=1), nn.GELU(), nn.Conv1d(d_model, d_model, 3, stride=2, padding=1), ) else: self.phase_proj = None # sharing schedule (v1.6-small/tiny): reuse physical layers per pattern. # enc_share/dec_share: int (#unique layers) or "E0,E1,E0,E1" label string. self.enc_share = enc_share self.dec_share = dec_share if (enc_share is None and dec_share is None and not adapter_rank and not adapter_rank_ffn and not depth_emb and not unshare_last_dec): enc_layer = nn.TransformerEncoderLayer( d_model, nhead, ffn, dropout, activation="gelu", batch_first=True, norm_first=True) self.encoder = nn.TransformerEncoder(enc_layer, enc_layers, nn.LayerNorm(d_model)) dec_layer = nn.TransformerDecoderLayer( d_model, nhead, ffn, dropout, activation="gelu", batch_first=True, norm_first=True) self.decoder = nn.TransformerDecoder(dec_layer, dec_layers, nn.LayerNorm(d_model)) else: enc_plan = _parse_pattern(enc_share, enc_layers) dec_plan = _parse_pattern(dec_share, dec_layers) # v1.7: un-share the LAST decoder layer (pattern-shaping layer in # front of the generation head gets its own physical parameters); # no-op if that depth is already unique. if unshare_last_dec and dec_plan.count(dec_plan[-1]) > 1: dec_plan = dec_plan[:-1] + [max(dec_plan) + 1] self.encoder = SharedEncoder(d_model, nhead, ffn, dropout, enc_plan, adapter_rank, unique_layernorm, adapter_rank_ffn, depth_emb) self.decoder = SharedDecoder(d_model, nhead, ffn, dropout, dec_plan, adapter_rank, unique_layernorm, adapter_rank_ffn, depth_emb) # ALBERT-style factorized embeddings (v1.5): vocab->E->d. Position # tokens dominate the table (1728/1810 rows), so E=64 cuts the # embedding block ~3.5x; logits stay tied via the projected matrix. self.emb_factor = emb_factor if emb_factor: self.tok_emb = nn.Embedding(vocab_size, emb_factor, padding_idx=VOCAB.pad) nn.init.normal_(self.tok_emb.weight, std=0.02) self.emb_proj = nn.Linear(emb_factor, d_model, bias=False) nn.init.normal_(self.emb_proj.weight, std=1.0 / (emb_factor ** 0.5) * 0.16) self.out = None else: self.tok_emb = nn.Embedding(vocab_size, d_model, padding_idx=VOCAB.pad) nn.init.normal_(self.tok_emb.weight, std=0.02) # keep tied logits well-scaled self.emb_proj = None self.out = nn.Linear(d_model, vocab_size, bias=False) self.out.weight = self.tok_emb.weight # weight tying self.dropout = nn.Dropout(dropout) # auxiliary per-position onset-heatmap head on the encoder (v2) self.aux = nn.Linear(d_model, 1) if aux else None # pointer alignment head: each generated note must point to its audio # frame (differentiable provenance; explainability that trains alignment) self.ptr = nn.Linear(d_model, d_model) if ptr else None # beat/downbeat head: explicit metrical percept on the encoder. # beat="hires" builds the frame-resolution head (11.6 ms bins) self.beat = (BeatHiRes(d_model) if beat == "hires" else (nn.Linear(d_model, 2) if beat else None)) # song-level context: coarse whole-song summary + window position (v4). # Lets the model see beyond the window, so intentional gaps (a breath # before a strong section) are informed decisions, not failures. # v2 fix (after the v4 negative result): summary chunks get DEDICATED # positional embeddings instead of reusing window positions. if global_ctx: self.gsum_proj = nn.Linear(N_MELS, d_model) self.seg_emb = nn.Parameter(torch.zeros(2, d_model)) self.pos_emb = nn.Embedding(16, d_model) # window position bucket self.gpos = nn.Parameter(torch.randn(128, d_model) * 0.02) else: self.gsum_proj = None # functional time embeddings: the 1728 TIME tokens share a sinusoidal # basis + small projection instead of free embeddings (fewer params, # neighbouring times get similar representations) if func_time: self.register_buffer("time_basis", sinusoidal(WINDOW, 64), persistent=False) self.time_proj = nn.Linear(64, d_model) # match the 0.02-std scale of tok_emb rows (basis row norm ~ sqrt(32)) nn.init.normal_(self.time_proj.weight, std=0.004) nn.init.zeros_(self.time_proj.bias) else: self.time_proj = None def encode(self, mel, gsum=None, pos_bucket=None): # mel: (B, C, T). clean_phase: C may be N_MELS (audio only) or N_MELS+2 # (audio + 2 grid-phase channels); only the audio channels reach the # encoder, so the returned memory is CLEAN (phase-free) and safe for the # beat head. gsum: (B, G, n_mels) whole-song summary chunks; # pos_bucket: (B,) window-position bucket in [0, 16). audio = mel[:, :N_MELS] if self.clean_phase else mel h = self.frontend(audio).transpose(1, 2) # (B, T/4, d) h = h + self.enc_pos[: h.shape[1]] if self.gsum_proj is not None and gsum is not None: g = self.gsum_proj(gsum) + self.seg_emb[1] + self.gpos[: gsum.shape[1]] h = h + self.seg_emb[0] p = self.pos_emb(pos_bucket).unsqueeze(1) # (B, 1, d) h = torch.cat([p, g, h], dim=1) return self.encoder(self.dropout(h)) def phase_mem(self, memory, mel): """Phase-conditioned copy of `memory` for the DECODER to cross-attend. Adds a projection of the grid-phase channels to the audio-frame slice of memory; the beat head keeps reading the clean `memory`. No-op unless clean_phase is on and phase channels are present. Time-mode windows (raw phase = -1) get valid=0 and zeroed phase -> ~zero conditioning.""" nch = N_MELS if not self.clean_phase or mel.shape[1] < nch + 2: return memory ph = mel[:, nch:nch + 2] # (B, 2, T) in [0,1); -1 marks no-grid frames valid = (ph[:, :1] >= 0).float() # 1 where a grid exists, else 0 ph_in = torch.cat([ph.clamp(min=0.0), valid], dim=1) # (B, 3, T) p = self.phase_proj(ph_in).transpose(1, 2) # (B, T/4, d) L = p.shape[1] out = memory.clone() out[:, -L:] = out[:, -L:] + p # audio frames are the last L memory slots return out def emb_matrix(self): """Full embedding matrix; factorized and/or functional rows resolved.""" E = self.tok_emb.weight if self.emb_proj is not None: E = self.emb_proj(E) if self.time_proj is not None: t = self.time_proj(self.time_basis) # (WINDOW, d) E = torch.cat([E[: VOCAB.time0], t, E[VOCAB.time0 + WINDOW :]], dim=0) return E def decode(self, tgt_in, memory): # tgt_in: (B, L) token ids L = tgt_in.shape[1] E = self.emb_matrix() h = nn.functional.embedding(tgt_in, E) * math.sqrt(self.d_model) + self.dec_pos[:L] mask = nn.Transformer.generate_square_subsequent_mask(L, device=tgt_in.device) pad_mask = tgt_in == VOCAB.pad h = self.decoder( self.dropout(h), memory, tgt_mask=mask, tgt_key_padding_mask=pad_mask, tgt_is_causal=True, ) return h @ E.T # tied output projection (functional rows included) def enable_ptr(self): self.ptr = nn.Linear(self.d_model, self.d_model) def enable_beat(self, hires=False): self.beat = BeatHiRes(self.d_model) if hires else nn.Linear(self.d_model, 2) def forward(self, mel, tgt, return_aux=False, gsum=None, pos_bucket=None, return_extras=False, in_tgt=None): # in_tgt: optional corrupted decoder input (e.g. MASKed types for the # skeleton->color infill curriculum); gold targets stay = tgt memory = self.encode(mel, gsum=gsum, pos_bucket=pos_bucket) # CLEAN dec_mem = self.phase_mem(memory, mel) # phase-conditioned for the decoder dec_in = (in_tgt if in_tgt is not None else tgt)[:, :-1] L = dec_in.shape[1] E = self.emb_matrix() import math as _m h = nn.functional.embedding(dec_in, E) * _m.sqrt(self.d_model) + self.dec_pos[:L] dm = nn.Transformer.generate_square_subsequent_mask(L, device=tgt.device) h = self.decoder(self.dropout(h), dec_mem, tgt_mask=dm, tgt_key_padding_mask=dec_in == VOCAB.pad, tgt_is_causal=True) logits = h @ E.T if not (return_aux or return_extras): return logits outs = [logits] # aux/ptr/beat heads read the CLEAN memory (no phase leak) aux_mem = memory[:, -(WINDOW // 4):] outs.append(self.aux(aux_mem).squeeze(-1) if self.aux is not None else None) if return_extras: ptr_logits = None if self.ptr is not None: q = self.ptr(h) # (B, L, d) ptr_logits = q @ aux_mem.transpose(1, 2) / _m.sqrt(self.d_model) beat_logits = self.beat(aux_mem) if self.beat is not None else None outs += [ptr_logits, beat_logits] return tuple(outs) def count_params(m): return sum(p.numel() for p in m.parameters() if p.requires_grad)