| 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 |
|
|
|
|
| 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) |
|
|
| |
| 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 |
|
|
| 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) |
| residual_d = h_d - h_d_hat |
| r_w = self.router_pr(residual_d) |
| 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 |
|
|