LFG-2 / deep_projector.py
glenn2's picture
Upload deep_projector.py with huggingface_hub
8b099fc verified
Raw
History Blame Contribute Delete
3.89 kB
"""Deeper audio projector for the HF Gemma4 unified model.
The stock model maps audio encoder features into LM space with a single
`nn.Linear(1536 -> 5376, bias=False)` at `model.model.embed_audio.embedding_projection`.
`DeepAudioProjector` is a drop-in replacement that keeps that linear as a
warm-started backbone and adds a *zero-initialised* residual MLP, so at init
the output is bit-identical to the original projector (we reuse the learned
phase-1 weights exactly), and SFT learns the deeper correction on top.
It exposes a `.weight` property because the parent `embed_audio.forward` reads
`self.embedding_projection.weight.dtype` to decide the input cast dtype.
"""
from __future__ import annotations
import torch
import torch.nn as nn
class DeepAudioProjector(nn.Module):
def __init__(
self,
in_dim: int = 1536,
out_dim: int = 5376,
hidden: int = 4096,
n_hidden_layers: int = 2,
dropout: float = 0.0,
out_dtype: torch.dtype = torch.bfloat16,
):
super().__init__()
self.in_dim = in_dim
self.out_dim = out_dim
self.out_dtype = out_dtype
# Warm-started backbone (copied from the trained phase-1 Linear).
self.proj = nn.Linear(in_dim, out_dim, bias=False)
# Residual branch (starts at exactly 0 -> identity behaviour at init).
self.ln = nn.LayerNorm(in_dim)
layers: list[nn.Module] = []
d = in_dim
for _ in range(n_hidden_layers):
layers += [nn.Linear(d, hidden), nn.GELU()]
if dropout > 0:
layers.append(nn.Dropout(dropout))
d = hidden
self.mlp = nn.Sequential(*layers)
self.out = nn.Linear(d, out_dim)
nn.init.zeros_(self.out.weight)
nn.init.zeros_(self.out.bias)
@property
def weight(self) -> torch.Tensor:
# Parent module reads `.weight.dtype`; expose the backbone weight.
return self.proj.weight
def forward(self, x: torch.Tensor) -> torch.Tensor:
wdt = self.proj.weight.dtype
x = x.to(wdt)
base = self.proj(x)
res = self.out(self.mlp(self.ln(x)))
return (base + res).to(self.out_dtype)
@classmethod
def from_linear(cls, linear: nn.Linear, **kwargs) -> "DeepAudioProjector":
out_dim, in_dim = linear.weight.shape
m = cls(in_dim=in_dim, out_dim=out_dim, **kwargs)
with torch.no_grad():
m.proj.weight.copy_(linear.weight)
return m
def find_audio_projection_parent(model):
"""Return (parent_module, attr_name) for the audio embedding_projection."""
inner = getattr(model, "model", model)
embed_audio = getattr(inner, "embed_audio", None)
if embed_audio is None:
raise AttributeError("Could not find model.model.embed_audio")
if not hasattr(embed_audio, "embedding_projection"):
raise AttributeError("embed_audio has no embedding_projection")
return embed_audio, "embedding_projection"
def install_deep_projector(
model,
hidden: int = 4096,
n_hidden_layers: int = 2,
dropout: float = 0.0,
param_dtype: torch.dtype = torch.float32,
):
"""Replace the audio embedding_projection with a warm-started DeepAudioProjector.
Returns the new projector module (params left in `param_dtype`, e.g. fp32 for
stable optimisation; forward output is cast back to the LM dtype).
"""
parent, attr = find_audio_projection_parent(model)
old = getattr(parent, attr)
assert isinstance(old, nn.Linear), f"expected nn.Linear, got {type(old)}"
out_dtype = old.weight.dtype
deep = DeepAudioProjector.from_linear(
old, hidden=hidden, n_hidden_layers=n_hidden_layers,
dropout=dropout, out_dtype=out_dtype,
)
deep = deep.to(device=old.weight.device, dtype=param_dtype)
setattr(parent, attr, deep)
return deep