""" Temporal Adapter (TA) from VocBulwark (Appendix B.1). Lightweight module injected at each upsample stage of the frozen vocoder. Embeds watermark bits into acoustic features via: 1. Acoustic Feature Alignment: Emb(w) → PFP → w_proj ∈ [B, C, 1] 2. Frame-level Temporal Broadcasting: w_proj → w_latent ∈ [B, C, T] 3. Adaptive Injection: concat(w_latent, h) → downsample → DSC → zero_conv + h (residual) Architecture from paper: Embedding: Linear(bits, 2*C) → LeakyReLU → Linear(2*C, 512) → LeakyReLU PFP: Linear(512, rank) → SiLU → Linear(rank, C) """ import torch import torch.nn as nn class TemporalAdapterBlock(nn.Module): """Single TA block for one upsample stage.""" def __init__(self, watermark_bits: int, channels: int, hidden_dim: int = 128, zero_conv_std: float = 0.02): super().__init__() self.channels = channels # Acoustic Feature Alignment (paper B.1): # Two FC layers with two LeakyReLU activations. # First FC projects to 2 * hidden_feature_channels, second to 512. self.embedding = nn.Sequential( nn.Linear(watermark_bits, channels * 2), nn.LeakyReLU(0.2), nn.Linear(channels * 2, 512), nn.LeakyReLU(0.2), ) # Progressive Feature Projection (PFP) (paper B.1): # Two FC layers with SiLU activation between them. # First FC reduces to pre-defined Rank, second aligns to C. rank = hidden_dim // 2 # pre-defined rank value self.pfp = nn.Sequential( nn.Linear(512, rank), nn.SiLU(), nn.Linear(rank, channels), ) # Adaptive Injection: concat [w_latent, h] → downsample → DSC → zero_conv # Downsample from 2C to C self.downsample = nn.Conv1d(channels * 2, channels, kernel_size=1) # Depth-wise Separable Convolution (DSC) # BatchNorm is load-bearing: it normalizes the DSC output to ~unit variance # before the small-init zero_conv, giving the injected watermark residual a # meaningful scale. Without any norm the residual is ~0, the watermark barely # perturbs the audio, extraction gradients vanish (grad_norm ~0.03) and # training sticks at exactly random (acc 0.5, loss_ext = ln 2). InstanceNorm # also fails (it cancels the time-constant watermark). So: keep BatchNorm. self.dsc = nn.Sequential( # Depth-wise conv nn.Conv1d(channels, channels, kernel_size=3, padding=1, groups=channels), nn.BatchNorm1d(channels), nn.LeakyReLU(0.2), # Point-wise conv nn.Conv1d(channels, channels, kernel_size=1), nn.BatchNorm1d(channels), nn.LeakyReLU(0.2), ) # Small-init convolution (not zero-init) to break the bootstrap deadlock. # Zero-init prevents any watermark signal from reaching the audio, # so the extractor can never learn. Small random init ensures different # watermarks produce slightly different audio from the start. self.zero_conv = nn.Conv1d(channels, channels, kernel_size=1) nn.init.normal_(self.zero_conv.weight, std=zero_conv_std) nn.init.zeros_(self.zero_conv.bias) def forward(self, h: torch.Tensor, watermark: torch.Tensor) -> torch.Tensor: """ Args: h: [B, C, T] hidden features from vocoder upsample stage watermark: [B, watermark_bits] binary watermark vector (float) Returns: [B, C, T] watermarked features (h + residual) """ B, C, T = h.shape # 1. Acoustic Feature Alignment w_emb = self.embedding(watermark) # [B, 512] w_proj = self.pfp(w_emb) # [B, C] # 2. Frame-level Temporal Broadcasting w_latent = w_proj.unsqueeze(-1).expand(B, C, T) # [B, C, T] # 3. Adaptive Injection h_cat = torch.cat([w_latent, h], dim=1) # [B, 2C, T] h_down = self.downsample(h_cat) # [B, C, T] h_dsc = self.dsc(h_down) # [B, C, T] residual = self.zero_conv(h_dsc) # [B, C, T] return h + residual