"""Self-contained PyTorch loader for ThomasThebaud/PROPS_K16.""" from __future__ import annotations import json from pathlib import Path import torch from torch import nn from torch.nn import functional as F class ComposedGMMMDN(nn.Module): """Map a 384-D description embedding to a GMM over 192-D speaker embeddings.""" def __init__( self, input_dim: int = 384, output_dim: int = 192, effective_components: int = 15072, hidden_dims=(1024, 2048, 1024), dropout: float = 0.1, min_sigma: float = 1e-4, ): super().__init__() self.input_dim = input_dim self.output_dim = output_dim self.K = effective_components self.min_sigma = min_sigma layers = [] previous = input_dim for width in hidden_dims: layers.extend( [nn.Linear(previous, width), nn.LayerNorm(width), nn.GELU(), nn.Dropout(dropout)] ) previous = width self.backbone = nn.Sequential(*layers) self.pi_head = nn.Linear(previous, effective_components) self.mu = nn.Parameter(torch.empty(effective_components, output_dim), requires_grad=False) self.raw_sigma = nn.Parameter( torch.empty(effective_components, output_dim), requires_grad=False ) # These heads are inherited, unused parameters in the original checkpoint. # Keeping them here permits an exact, strict state-dict load. requested_components = 16 self.mu_head = nn.Linear(previous, requested_components * output_dim) self.sigma_head = nn.Linear(previous, requested_components * output_dim) def forward(self, description_embeddings: torch.Tensor): h = self.backbone(description_embeddings) pi_logits = self.pi_head(h) sigma = F.softplus(self.raw_sigma) + self.min_sigma return pi_logits, self.mu.unsqueeze(0).expand(h.shape[0], -1, -1), sigma.unsqueeze(0).expand( h.shape[0], -1, -1 ) @torch.inference_mode() def distribution(self, description_embeddings: torch.Tensor): """Return normalized weights, component means, and diagonal standard deviations.""" pi_logits, mu, sigma = self(description_embeddings) return torch.softmax(pi_logits, dim=-1), mu, sigma @torch.inference_mode() def sample(self, description_embeddings: torch.Tensor, samples_per_prompt: int = 1): """Draw speaker embeddings; output shape is [batch, samples_per_prompt, 192].""" pi, mu, sigma = self.distribution(description_embeddings) component = torch.multinomial(pi, samples_per_prompt, replacement=True) gather_index = component.unsqueeze(-1).expand(-1, -1, self.output_dim) chosen_mu = mu.gather(1, gather_index) chosen_sigma = sigma.gather(1, gather_index) return chosen_mu + chosen_sigma * torch.randn_like(chosen_mu) @classmethod def from_pretrained(cls, repo_or_path="ThomasThebaud/PROPS_K16", device="cpu"): path = Path(repo_or_path) if path.is_dir(): config_path = path / "config.json" else: from huggingface_hub import hf_hub_download config_path = Path(hf_hub_download(repo_or_path, "config.json")) config = json.loads(config_path.read_text()) if path.is_dir(): checkpoint_path = path / config["checkpoint_file"] else: checkpoint_path = Path(hf_hub_download(repo_or_path, config["checkpoint_file"])) model = cls( input_dim=config["input_dim"], output_dim=config["output_dim"], effective_components=config["effective_components"], hidden_dims=tuple(config["hidden_dims"]), dropout=config["dropout"], min_sigma=config["min_sigma"], ) try: checkpoint = torch.load(checkpoint_path, map_location="cpu", weights_only=True) except TypeError: checkpoint = torch.load(checkpoint_path, map_location="cpu") model.load_state_dict(checkpoint["model_state_dict"], strict=True) return model.to(device).eval()