Spaces:
Running on Zero
Running on Zero
File size: 6,565 Bytes
fdcef0f | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 | from __future__ import annotations
from pathlib import Path
import torch
import torch.nn as nn
from safetensors.torch import load_file as load_safetensors_file
from safetensors.torch import save_file as save_safetensors_file
SPEAKER_INVERSION_UNCOND_MODES = {"mask", "noise"}
SPEAKER_INVERSION_SAFETENSORS_SUFFIX = ".speaker.safetensors"
SPEAKER_EMBEDDING_KEY = "speaker_embedding"
def normalize_speaker_embedding_tensor(
tensor: torch.Tensor,
*,
speaker_dim: int,
field_name: str = SPEAKER_EMBEDDING_KEY,
) -> torch.Tensor:
if tensor.ndim == 3 and tensor.shape[0] == 1:
tensor = tensor[0]
if tensor.ndim != 2:
raise ValueError(f"{field_name} must have shape (tokens, dim), got {tuple(tensor.shape)}")
if int(tensor.shape[0]) <= 0:
raise ValueError(f"{field_name} must contain at least one token.")
if int(tensor.shape[1]) != int(speaker_dim):
raise ValueError(
f"{field_name} dim mismatch: expected {int(speaker_dim)}, got {int(tensor.shape[1])}"
)
return tensor.detach().float().contiguous()
def is_speaker_inversion_safetensors_path(path: str | Path) -> bool:
return Path(path).name.endswith(SPEAKER_INVERSION_SAFETENSORS_SUFFIX)
class SpeakerInversionEmbedding(nn.Module):
"""Learned speaker/style tokens that bypass the reference latent speaker encoder."""
def __init__(
self,
*,
num_tokens: int,
speaker_dim: int,
init_std: float,
init_embedding: torch.Tensor | None = None,
) -> None:
super().__init__()
num_tokens = int(num_tokens)
speaker_dim = int(speaker_dim)
init_std = float(init_std)
if num_tokens <= 0:
raise ValueError(f"speaker inversion tokens must be > 0, got {num_tokens}")
if speaker_dim <= 0:
raise ValueError(f"speaker_dim must be > 0, got {speaker_dim}")
if init_std < 0:
raise ValueError(f"speaker inversion init_std must be >= 0, got {init_std}")
if init_embedding is None:
embedding = torch.randn(num_tokens, speaker_dim, dtype=torch.float32) * init_std
else:
embedding = normalize_speaker_embedding_tensor(
init_embedding,
speaker_dim=speaker_dim,
field_name=SPEAKER_EMBEDDING_KEY,
)
if int(embedding.shape[0]) != num_tokens:
raise ValueError(
"speaker inversion init embedding token mismatch: "
f"expected {num_tokens}, got {int(embedding.shape[0])}"
)
self.embedding = nn.Parameter(embedding)
@property
def num_tokens(self) -> int:
return int(self.embedding.shape[0])
@property
def speaker_dim(self) -> int:
return int(self.embedding.shape[1])
def forward(
self,
*,
batch_size: int,
device: torch.device,
dtype: torch.dtype,
) -> tuple[torch.Tensor, torch.Tensor]:
state = self.embedding.to(device=device, dtype=dtype)[None, :, :].expand(
int(batch_size),
-1,
-1,
)
mask = torch.ones((int(batch_size), self.num_tokens), dtype=torch.bool, device=device)
return state, mask
def _extract_embedding_payload(raw: dict[str, torch.Tensor]) -> dict[str, torch.Tensor]:
if not isinstance(raw, dict):
raise ValueError(
f"Speaker inversion file must contain a tensor dictionary, got {type(raw)!r}."
)
if SPEAKER_EMBEDDING_KEY in raw:
embedding = raw[SPEAKER_EMBEDDING_KEY]
if not isinstance(embedding, torch.Tensor):
raise ValueError(
f"Speaker inversion '{SPEAKER_EMBEDDING_KEY}' must be a tensor, "
f"got {type(embedding)!r}."
)
return {SPEAKER_EMBEDDING_KEY: embedding}
raise ValueError(f"Speaker inversion file is missing '{SPEAKER_EMBEDDING_KEY}'.")
def normalize_speaker_inversion_payload(
raw: dict[str, torch.Tensor],
) -> dict[str, torch.Tensor]:
payload = _extract_embedding_payload(raw)
embedding = payload[SPEAKER_EMBEDDING_KEY]
out: dict[str, torch.Tensor] = {
SPEAKER_EMBEDDING_KEY: embedding,
}
return out
def load_speaker_inversion_payload(
path: str | Path,
) -> dict[str, torch.Tensor]:
source = Path(path).expanduser()
if not is_speaker_inversion_safetensors_path(source):
raise ValueError(
"Speaker Inversion embeddings must use the "
f"{SPEAKER_INVERSION_SAFETENSORS_SUFFIX!r} suffix: {source}"
)
raw = load_safetensors_file(source, device="cpu")
out = normalize_speaker_inversion_payload(raw)
return out
def save_speaker_inversion_safetensors(
path: str | Path,
payload: dict[str, torch.Tensor],
*,
dtype: torch.dtype = torch.float32,
) -> None:
target = Path(path)
if not is_speaker_inversion_safetensors_path(target):
raise ValueError(
"Speaker Inversion safetensors output must use the "
f"{SPEAKER_INVERSION_SAFETENSORS_SUFFIX!r} suffix: {target}"
)
normalized = normalize_speaker_inversion_payload(payload)
tensors = {
SPEAKER_EMBEDDING_KEY: normalized[SPEAKER_EMBEDDING_KEY].to(dtype=dtype),
}
target.parent.mkdir(parents=True, exist_ok=True)
save_safetensors_file(tensors, str(target), metadata={})
def speaker_inversion_batch_tensors(
speaker_embedding: torch.Tensor,
batch_size: int,
device: torch.device,
dtype: torch.dtype,
) -> tuple[torch.Tensor, torch.Tensor]:
embedding = speaker_embedding.to(device=device, dtype=dtype)
state = embedding[None, :, :].expand(int(batch_size), -1, -1)
mask = torch.ones((int(batch_size), embedding.shape[0]), dtype=torch.bool, device=device)
return state, mask
def speaker_inversion_state_dict(model: nn.Module) -> dict[str, torch.Tensor]:
module = getattr(model, "speaker_inversion", None)
if not isinstance(module, SpeakerInversionEmbedding):
raise ValueError("Model does not have an enabled SpeakerInversionEmbedding module.")
return {
SPEAKER_EMBEDDING_KEY: module.embedding.detach().cpu().float().clone(),
}
def save_speaker_inversion_checkpoint(
path: str | Path,
*,
model: nn.Module,
) -> None:
path = Path(path)
path.parent.mkdir(parents=True, exist_ok=True)
state = speaker_inversion_state_dict(model)
save_speaker_inversion_safetensors(path, state)
|