JacobLinCool's picture
Restore selectable legacy model generations
be60cf4
Raw
History Blame Contribute Delete
20.9 kB
"""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)