PROPS_K16 / model.py
ThomasThebaud's picture
Add PROPS K16 checkpoint, model card, and inference code
a4534ab verified
Raw History Blame Contribute Delete
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
)
@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()