implicit_fusion / encoder.py
yucongzh's picture
Upload folder using huggingface_hub
704bc5d verified
Raw
History Blame Contribute Delete
16.6 kB
import sys
from pathlib import Path
import torch
import torch.nn as nn
import torch.nn.functional as F
from transformers import AutoFeatureExtractor, AutoModel, WhisperModel, WhisperProcessor
_THIS_DIR = Path(__file__).resolve().parent
if str(_THIS_DIR) not in sys.path:
sys.path.insert(0, str(_THIS_DIR))
from _checkpoint_mixin import load_audio_encoder_checkpoint # noqa: E402
def length_to_mask(lengths: torch.Tensor, max_len: int | None = None) -> torch.Tensor:
if max_len is None:
max_len = int(lengths.max().item())
idx = torch.arange(max_len, device=lengths.device).unsqueeze(0)
return (idx < lengths.unsqueeze(1)).long()
def _audio_lengths(audio: torch.Tensor, audio_attention_mask: torch.Tensor | None) -> torch.Tensor:
if audio_attention_mask is None:
return torch.full((audio.shape[0],), audio.shape[-1], device=audio.device, dtype=torch.long)
return audio_attention_mask.sum(-1).to(torch.long)
def _ceil_div(a: torch.Tensor, b: int) -> torch.Tensor:
return (a + (b - 1)) // b
def _resample_mask(mask: torch.Tensor, target_len: int) -> torch.Tensor:
if mask.shape[1] == target_len:
return mask
out = F.interpolate(mask.float().unsqueeze(1), size=target_len, mode="nearest")
return out.squeeze(1).to(torch.long)
def _align_time(x: torch.Tensor, target_len: int) -> torch.Tensor:
if x.shape[1] == target_len:
return x
x_t = x.transpose(1, 2)
x_t = F.interpolate(x_t, size=target_len, mode="nearest")
return x_t.transpose(1, 2)
class DKU_WHU2_Encoder(torch.nn.Module):
"""Softmax2 token gate + log-mel/STFT residual."""
def __init__(
self,
dasheng_model_name: str = "mispeech/dasheng-base",
whisper_model_name: str = "openai/whisper-base",
gate_beta: float = 0.1,
gate_align: str = "project",
gate_target_dim: int | None = None,
target_len: str = "min",
temperature: float = 1.0,
init_backbones: bool = True,
mock_dasheng_dim: int = 768,
mock_whisper_dim: int = 768,
n_fft: int = 1024,
hop_length: int = 640,
n_mels: int = 128,
use_residual_aux: bool = False,
w_recon: float = 1.0,
w_decor: float = 1e-3,
checkpoint_dir: str | None = "checkpoint-10000",
checkpoint_step: int | None = 50000,
load_pretrained: bool = True,
load_strict: bool = False,
router_only: bool = False,
verbose: bool = True,
) -> None:
super().__init__()
self.sampling_rate = 16000
self.hop_size_in_ms = 40
self._target_len = str(target_len)
self._gate_align = str(gate_align)
self._temperature = float(temperature)
self._use_residual_aux = bool(use_residual_aux)
self.router_w_recon = float(w_recon)
self.router_w_decor = float(w_decor)
self._init_backbones = bool(init_backbones)
if self._init_backbones:
self._dasheng_processor = AutoFeatureExtractor.from_pretrained(dasheng_model_name, trust_remote_code=True)
self._dasheng_model = AutoModel.from_pretrained(
dasheng_model_name,
trust_remote_code=True,
low_cpu_mem_usage=False,
device_map=None,
)
self._whisper_processor = WhisperProcessor.from_pretrained(whisper_model_name)
self._whisper_encoder = WhisperModel.from_pretrained(
whisper_model_name,
low_cpu_mem_usage=False,
device_map=None,
).get_encoder()
dasheng_dim = self._infer_dasheng_dim()
whisper_dim = int(self._whisper_encoder.config.d_model)
else:
self._dasheng_processor = None
self._dasheng_model = None
self._whisper_processor = None
self._whisper_encoder = None
dasheng_dim = int(mock_dasheng_dim)
whisper_dim = int(mock_whisper_dim)
self._dasheng_dim = int(dasheng_dim)
self._whisper_dim = int(whisper_dim)
target_dim = int(gate_target_dim) if gate_target_dim is not None else int(min(dasheng_dim, whisper_dim))
if target_dim <= 0:
raise ValueError("gate_target_dim must be positive")
self._gate_dim = target_dim
self.register_buffer("router_beta", torch.tensor(float(gate_beta), dtype=torch.float32), persistent=True)
self.router_proj_dasheng: nn.Module | None = None
self.router_proj_whisper: nn.Module | None = None
if self._gate_align not in ("truncate", "project"):
raise ValueError("gate_align must be one of: truncate, project")
if self._gate_align == "project":
if int(dasheng_dim) != target_dim:
self.router_proj_dasheng = nn.Linear(int(dasheng_dim), target_dim, bias=False)
if int(whisper_dim) != target_dim:
self.router_proj_whisper = nn.Linear(int(whisper_dim), target_dim, bias=False)
self.router_gate_linear = nn.Linear(2 * target_dim, 2)
self.router_fusion_ln = nn.LayerNorm(target_dim)
self.router_sum_ln = nn.LayerNorm(target_dim)
# Optional residual-orthogonalized aux branch.
self.router_w2d: nn.Module | None = None
self.router_pr: nn.Module | None = None
if self._use_residual_aux:
self.router_w2d = nn.Linear(self._whisper_dim, self._dasheng_dim, bias=False)
self.router_pr = nn.Linear(self._dasheng_dim, self._whisper_dim)
self._n_fft = int(n_fft)
self._hop_length = int(hop_length)
self._n_mels = int(n_mels)
self.router_spec_proj = nn.Linear(self._n_mels, target_dim)
self.router_spec_alpha = nn.Parameter(torch.zeros(1))
self.router_spec_ln = nn.LayerNorm(target_dim)
self.output_dim = target_dim
self.aux_loss = torch.tensor(0.0)
self.router_reg_loss = torch.tensor(0.0)
self.recon_loss = torch.tensor(0.0)
self.decor_loss = torch.tensor(0.0)
self.aux_loss_total = torch.tensor(0.0)
self.aux_items: dict[str, torch.Tensor] = {}
mel_fb = self._build_mel_filterbank(self._n_fft, self._n_mels, self.sampling_rate)
if mel_fb is not None:
self.register_buffer("mel_fb", mel_fb, persistent=False)
else:
self.mel_fb = None
if load_pretrained:
load_audio_encoder_checkpoint(
self,
checkpoint_dir,
checkpoint_step=checkpoint_step,
strict=load_strict,
router_only=router_only,
verbose=verbose,
)
def _build_mel_filterbank(self, n_fft: int, n_mels: int, sr: int) -> torch.Tensor | None:
try:
import torchaudio.functional as AF # type: ignore
fb = AF.melscale_fbanks(
n_freqs=n_fft // 2 + 1,
f_min=0.0,
f_max=float(sr // 2),
n_mels=n_mels,
sample_rate=sr,
)
return fb
except Exception:
return None
def _infer_dasheng_dim(self) -> int:
cfg = getattr(self._dasheng_model, "config", None)
if cfg is None:
raise ValueError("Dasheng model has no config; cannot infer output_dim")
for attr in ("hidden_size", "d_model", "embed_dim"):
if hasattr(cfg, attr):
return int(getattr(cfg, attr))
if hasattr(cfg, "encoder_kwargs") and isinstance(cfg.encoder_kwargs, dict) and "embed_dim" in cfg.encoder_kwargs:
return int(cfg.encoder_kwargs["embed_dim"])
if hasattr(cfg, "encoder_kwargs") and isinstance(cfg.encoder_kwargs, dict) and "d_model" in cfg.encoder_kwargs:
return int(cfg.encoder_kwargs["d_model"])
raise ValueError("Could not infer Dasheng embedding dim from config")
def _dasheng_forward(
self, audio: torch.Tensor, audio_attention_mask: torch.Tensor | None
) -> tuple[torch.Tensor, torch.Tensor]:
if self._dasheng_model is None or self._dasheng_processor is None:
raise RuntimeError("Dasheng backbone not initialized.")
features = self._dasheng_processor(audio, return_tensors="pt")
model_device = next(self._dasheng_model.parameters()).device
features = {k: v.to(model_device) if isinstance(v, torch.Tensor) else v for k, v in features.items()}
out = self._dasheng_model(**features)
if hasattr(out, "last_hidden_state") and isinstance(out.last_hidden_state, torch.Tensor):
feats = out.last_hidden_state
elif hasattr(out, "hidden_states"):
hs = out.hidden_states
feats = hs[-1] if isinstance(hs, (tuple, list)) else hs
elif isinstance(out, (tuple, list)) and len(out) > 0 and isinstance(out[0], torch.Tensor):
feats = out[0]
else:
raise ValueError("Unexpected Dasheng model output; cannot get features")
t1 = int(feats.shape[1])
hop_samples = int(self.sampling_rate * self.hop_size_in_ms / 1000)
lengths = _audio_lengths(
audio.to(feats.device),
audio_attention_mask.to(feats.device) if audio_attention_mask is not None else None,
)
feat_lens = torch.clamp(_ceil_div(lengths, hop_samples), min=1, max=t1)
mask = length_to_mask(feat_lens, max_len=t1).to(feats.device)
return feats, mask
def _whisper_forward(
self, audio: torch.Tensor, audio_attention_mask: torch.Tensor | None
) -> tuple[torch.Tensor, torch.Tensor]:
if self._whisper_encoder is None or self._whisper_processor is None:
raise RuntimeError("Whisper backbone not initialized.")
audio_list = [a.detach().cpu().numpy() for a in audio]
if audio_attention_mask is None:
audio_lens = torch.tensor([a.shape[-1] for a in audio_list], dtype=torch.long)
else:
audio_lens = audio_attention_mask.sum(-1).detach().cpu().to(torch.long)
hop_samples = int(self.sampling_rate * self.hop_size_in_ms / 1000)
feature_lengths = torch.clamp(_ceil_div(audio_lens, hop_samples), min=1)
trim_length = int(feature_lengths.max().item())
attention_mask = length_to_mask(feature_lengths, max_len=trim_length)
feats = self._whisper_processor(audio_list, sampling_rate=self.sampling_rate, return_tensors="pt")
model_device = next(self._whisper_encoder.parameters()).device
feats = {k: v.to(model_device) if isinstance(v, torch.Tensor) else v for k, v in feats.items()}
out = self._whisper_encoder(**feats).last_hidden_state
out = out[:, :trim_length, :]
attention_mask = attention_mask[:, : out.shape[1]].to(out.device)
return out, attention_mask
def _align_gate_dims(self, d: torch.Tensor, w: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
if self._gate_align == "project":
if self.router_proj_dasheng is not None:
d = self.router_proj_dasheng(d)
if self.router_proj_whisper is not None:
w = self.router_proj_whisper(w)
return d, w
target = int(self._gate_dim)
return d[..., :target], w[..., :target]
def _softmax2_fuse(self, d: torch.Tensor, w: torch.Tensor) -> torch.Tensor:
gate_inp = torch.cat([d, w], dim=-1)
logits = self.router_gate_linear(gate_inp) / self._temperature
weights = torch.softmax(logits, dim=-1)
scaled = weights * self.router_beta
w_d = scaled[..., 0:1]
w_w = scaled[..., 1:2]
fused = w_d * d + w_w * w
fused = self.router_fusion_ln(fused)
fused = self.router_sum_ln(fused + d + w)
return fused
def _compute_log_mel(self, audio: torch.Tensor) -> torch.Tensor:
window = torch.hann_window(self._n_fft, device=audio.device)
spec = torch.stft(
audio,
n_fft=self._n_fft,
hop_length=self._hop_length,
win_length=self._n_fft,
window=window,
center=True,
return_complex=True,
)
mag = spec.abs() ** 2
if self.mel_fb is not None:
mel = torch.matmul(mag.transpose(1, 2), self.mel_fb.to(mag.device)).transpose(1, 2)
else:
mel = F.interpolate(mag.unsqueeze(1), size=(self._n_mels, mag.shape[2]), mode="nearest").squeeze(1)
mel = mel.clamp_min(1e-10).log()
return mel.transpose(1, 2)
def _masked_mean(self, values: torch.Tensor, mask: torch.Tensor) -> torch.Tensor:
weights = mask.to(dtype=values.dtype)
denom = weights.sum().clamp_min(1.0)
return (values * weights).sum() / denom
def _cross_cov_decorrelation_loss(
self,
h_w: torch.Tensor,
r_w: torch.Tensor,
mask: torch.Tensor,
) -> torch.Tensor:
x = h_w.reshape(-1, h_w.shape[-1])
y = r_w.reshape(-1, r_w.shape[-1])
keep = mask.reshape(-1) > 0
if keep.any():
x = x[keep]
y = y[keep]
n = int(x.shape[0])
if n <= 1:
return h_w.new_zeros((), dtype=torch.float32)
x = x.float()
y = y.float()
x_centered = x - x.mean(dim=0, keepdim=True)
y_centered = y - y.mean(dim=0, keepdim=True)
cov = torch.matmul(x_centered.transpose(0, 1), y_centered) / float(max(n - 1, 1))
return cov.pow(2).mean()
def _compute_residual_aux(
self,
h_d: torch.Tensor,
h_w: torch.Tensor,
mask: torch.Tensor,
) -> tuple[torch.Tensor, torch.Tensor]:
if self.router_w2d is None or self.router_pr is None:
return h_w.new_zeros((), dtype=torch.float32), h_w.new_zeros((), dtype=torch.float32)
h_d_hat = self.router_w2d(h_w) # [B,T,768]
residual_d = h_d - h_d_hat
r_w = self.router_pr(residual_d) # [B,T,512]
recon_token = residual_d.float().pow(2).mean(dim=-1)
recon_loss = self._masked_mean(recon_token, mask)
decor_loss = self._cross_cov_decorrelation_loss(h_w, r_w, mask)
return recon_loss, decor_loss
def forward(
self,
audio: torch.Tensor,
audio_attention_mask: torch.Tensor | None = None,
) -> tuple[torch.Tensor, torch.Tensor]:
if not isinstance(audio, torch.Tensor):
raise TypeError("audio must be a torch.Tensor")
if audio.ndim != 2:
raise ValueError("Expected audio shape [B, T]")
dasheng_feat, dasheng_mask = self._dasheng_forward(audio, audio_attention_mask)
whisper_feat, whisper_mask = self._whisper_forward(audio, audio_attention_mask)
if self._target_len == "whisper":
target_len = int(whisper_feat.shape[1])
elif self._target_len == "dasheng":
target_len = int(dasheng_feat.shape[1])
else:
target_len = int(min(dasheng_feat.shape[1], whisper_feat.shape[1]))
dasheng_mask = _resample_mask(dasheng_mask, target_len)
whisper_mask = _resample_mask(whisper_mask, target_len)
mask = (dasheng_mask & whisper_mask).long()
dasheng_feat = dasheng_feat[:, :target_len, :]
whisper_feat = whisper_feat[:, :target_len, :]
h_d_raw = dasheng_feat
h_w_raw = whisper_feat
dasheng_feat, whisper_feat = self._align_gate_dims(dasheng_feat, whisper_feat)
fused = self._softmax2_fuse(dasheng_feat, whisper_feat)
spec = self._compute_log_mel(audio)
if spec.ndim != 3:
raise ValueError("spec must be [B,T,F]")
if spec.shape[1] != fused.shape[1]:
spec = _align_time(spec, fused.shape[1])
spec_proj = self.router_spec_proj(spec)
spec_proj = self.router_spec_ln(spec_proj)
fused = fused + self.router_spec_alpha * spec_proj
recon_loss, decor_loss = self._compute_residual_aux(h_d_raw, h_w_raw, mask)
decor_w = float(getattr(self, "router_w_decor_runtime", self.router_w_decor))
aux_loss_total = fused.new_zeros((), dtype=torch.float32)
if self._use_residual_aux:
aux_loss_total = (self.router_w_recon * recon_loss) + (decor_w * decor_loss)
self.recon_loss = recon_loss
self.decor_loss = decor_loss
self.aux_loss_total = aux_loss_total
self.aux_items = {
"recon_loss": recon_loss,
"decor_loss": decor_loss,
"aux_loss_total": aux_loss_total,
}
return fused, mask