Download model.py from ThomasThebaud/PROPS_K16: direct link, hf CLI and curl.
- Browser
- Download file 4.19 kB
-
https://huggingface.co/ThomasThebaud/PROPS_K16/resolve/main/model.py
- Command line
-
hf download hf://ThomasThebaud/PROPS_K16/model.py
-
curl -L -o model.py https://huggingface.co/ThomasThebaud/PROPS_K16/resolve/main/model.py
4.19 kB
| """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 | |
| ) | |
| 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 | |
| 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) | |
| 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() | |