Files
Aletheia/tools/tts/booktts_student.py

436 lines
22 KiB
Python

#!/usr/bin/env python3
"""Build and train a compact drop-in BERT student for Aletheia BookTTS."""
from __future__ import annotations
import argparse
import json
import random
import shutil
import sys
import time
from collections import Counter
from pathlib import Path
import numpy as np
from vosk_booktts_experiment import WordPieceTokenizer, load_corpus, normalize_text
def build_dataset(args: argparse.Namespace) -> None:
import onnxruntime as ort
args.output.mkdir(parents=True, exist_ok=False)
tokenizer = WordPieceTokenizer(args.assets / "vocab.txt")
options = ort.SessionOptions()
options.log_severity_level = 3
teacher = ort.InferenceSession(
str(args.assets / "bert.int8.onnx"),
sess_options=options,
providers=["CPUExecutionProvider"],
)
token_counts: Counter[int] = Counter()
pending: list[tuple[str, np.ndarray]] = []
shard_samples: list[tuple[np.ndarray, np.ndarray]] = []
shard_records = []
sample_index = 0
shard_index = 0
teacher_dtype = np.float16 if args.teacher_dtype == "float16" else np.float32
def flush_shard() -> None:
nonlocal shard_index
if not shard_samples:
return
longest = max(ids.size for ids, _ in shard_samples)
shard_ids = np.zeros((len(shard_samples), longest), dtype=np.int64)
shard_mask = np.zeros_like(shard_ids)
shard_teacher = np.zeros((len(shard_samples), longest, 768), dtype=teacher_dtype)
for row, (ids, embeddings) in enumerate(shard_samples):
shard_ids[row, : ids.size] = ids
shard_mask[row, : ids.size] = 1
shard_teacher[row, : ids.size] = embeddings
filename = f"shard-{shard_index:05d}.npz"
np.savez(
args.output / filename,
original_input_ids=shard_ids,
attention_mask=shard_mask,
teacher=shard_teacher,
)
shard_records.append({"file": filename, "samples": len(shard_samples), "max_tokens": longest})
shard_samples.clear()
shard_index += 1
def flush_batch(manifest) -> None:
nonlocal sample_index
if not pending:
return
longest = max(ids.size for _, ids in pending)
batch_ids = np.zeros((len(pending), longest), dtype=np.int64)
batch_mask = np.zeros_like(batch_ids)
for row, (_, ids) in enumerate(pending):
batch_ids[row, : ids.size] = ids
batch_mask[row, : ids.size] = 1
outputs = teacher.run(None, {
"input_ids": batch_ids,
"attention_mask": batch_mask,
"token_type_ids": np.zeros_like(batch_ids),
})[0]
if outputs.ndim == 2:
outputs = outputs[np.newaxis, :, :]
for row, (text, ids) in enumerate(pending):
embeddings = outputs[row, : ids.size].astype(teacher_dtype, copy=False)
if embeddings.shape != (ids.size, 768):
raise ValueError(f"Unexpected teacher shape {embeddings.shape}")
shard_samples.append((ids.copy(), embeddings.copy()))
manifest.write(json.dumps({"sample": sample_index, "text": text, "tokens": int(ids.size)}, ensure_ascii=False) + "\n")
sample_index += 1
if len(shard_samples) >= args.shard_size:
flush_shard()
pending.clear()
with (args.output / "manifest.jsonl").open("w", encoding="utf-8", newline="\n") as manifest:
for text in load_corpus(args.corpus):
ids, _ = tokenizer.encode(normalize_text(text))
token_counts.update(int(value) for value in ids)
pending.append((text, ids))
if len(pending) >= args.teacher_batch_size:
flush_batch(manifest)
if sample_index % 1_000 == 0:
print(f"teacher_samples={sample_index}", flush=True)
flush_batch(manifest)
flush_shard()
observed = set(token_counts)
if args.reuse_token_map:
source_metadata = json.loads((args.reuse_token_map / "dataset.json").read_text(encoding="utf-8"))
source_map = np.load(args.reuse_token_map / "token_id_map.npy")
if source_map.size != len(tokenizer.tokens):
raise ValueError("Reused token map and teacher vocabulary have different sizes")
shutil.copy2(args.reuse_token_map / "token_id_map.npy", args.output / "token_id_map.npy")
shutil.copy2(args.reuse_token_map / "vocab.txt", args.output / "vocab.txt")
vocab_size = int(source_metadata["vocab_size"])
unknown_new = int(source_map[tokenizer.vocabulary["[UNK]"]])
replaced_occurrences = sum(
count for token, count in token_counts.items()
if token != tokenizer.vocabulary["[UNK]"] and int(source_map[token]) == unknown_new
)
else:
special_ids = [tokenizer.vocabulary[name] for name in ("[PAD]", "[UNK]", "[CLS]", "[SEP]")]
if args.max_vocab and len(observed) > args.max_vocab:
retained = set(special_ids)
retained.update(token for token, _ in token_counts.most_common(args.max_vocab - len(retained)))
else:
retained = observed | set(special_ids)
ordered_old_ids = special_ids + sorted(retained - set(special_ids))
old_to_new = {old: new for new, old in enumerate(ordered_old_ids)}
unknown_new = old_to_new[tokenizer.vocabulary["[UNK]"]]
reduced_tokens = [tokenizer.tokens[index] for index in ordered_old_ids]
(args.output / "vocab.txt").write_text("\n".join(reduced_tokens) + "\n", encoding="utf-8")
token_id_map = np.full(len(tokenizer.tokens), unknown_new, dtype=np.int64)
for old, new in old_to_new.items():
token_id_map[old] = new
np.save(args.output / "token_id_map.npy", token_id_map)
vocab_size = len(reduced_tokens)
replaced_occurrences = sum(count for token, count in token_counts.items() if token not in retained)
metadata = {
"samples": sample_index,
"vocab_size": vocab_size,
"observed_original_tokens": len(observed),
"replaced_token_occurrences": replaced_occurrences,
"reused_token_map": str(args.reuse_token_map) if args.reuse_token_map else None,
"teacher_dtype": args.teacher_dtype,
"teacher_batch_size": args.teacher_batch_size,
"shard_size": args.shard_size,
"shards": shard_records,
}
(args.output / "dataset.json").write_text(json.dumps(metadata, ensure_ascii=False, indent=2) + "\n", encoding="utf-8")
print(json.dumps(metadata, indent=2))
def train_student(args: argparse.Namespace) -> None:
import torch
from torch import nn
from torch.utils.data import DataLoader, Dataset, IterableDataset, Subset
class DistillationDataset(Dataset):
def __init__(self, root: Path):
self.files = sorted(root.glob("sample-*.npz"))
if not self.files:
raise ValueError(f"No samples in {root}")
mapping = root / "token_id_map.npy"
self.token_id_map = np.load(mapping) if mapping.exists() else None
def __len__(self) -> int:
return len(self.files)
def __getitem__(self, index: int) -> tuple[torch.Tensor, torch.Tensor]:
with np.load(self.files[index]) as data:
if "original_input_ids" in data:
original = data["original_input_ids"]
if self.token_id_map is None:
raise ValueError("token_id_map.npy is required for original_input_ids")
input_ids = self.token_id_map[original]
else:
input_ids = data["input_ids"]
teacher = data["teacher"].astype(np.float32)
return torch.from_numpy(input_ids.copy()), torch.from_numpy(teacher.copy())
class ShardedDataset(IterableDataset):
def __init__(self, root: Path, records: list[dict], shuffle: bool, seed: int):
self.root = root
self.records = list(records)
self.shuffle = shuffle
self.seed = seed
self.epoch = 0
self.token_id_map = np.load(root / "token_id_map.npy")
def __len__(self) -> int:
return sum(record["samples"] for record in self.records)
def set_epoch(self, epoch: int) -> None:
self.epoch = epoch
def __iter__(self):
rng = random.Random(self.seed + self.epoch)
records = list(self.records)
if self.shuffle:
rng.shuffle(records)
for record in records:
with np.load(self.root / record["file"]) as data:
original = data["original_input_ids"]
masks = data["attention_mask"]
teachers = data["teacher"]
rows = list(range(original.shape[0]))
if self.shuffle:
rng.shuffle(rows)
for row in rows:
length = int(masks[row].sum())
ids = self.token_id_map[original[row, :length]]
teacher = teachers[row, :length].astype(np.float32)
yield torch.from_numpy(ids.copy()), torch.from_numpy(teacher.copy())
def collate(batch):
longest = max(ids.size(0) for ids, _ in batch)
ids = torch.zeros((len(batch), longest), dtype=torch.long)
mask = torch.zeros((len(batch), longest), dtype=torch.long)
targets = torch.zeros((len(batch), longest, 768), dtype=torch.float32)
for row, (sample_ids, teacher) in enumerate(batch):
length = sample_ids.size(0)
ids[row, :length] = sample_ids
mask[row, :length] = 1
targets[row, :length] = teacher
return ids, mask, targets
class Student(nn.Module):
def __init__(self, vocab_size: int):
super().__init__()
self.token_embedding = nn.Embedding(vocab_size, args.hidden, padding_idx=0)
self.position_embedding = nn.Embedding(args.max_length, args.hidden)
self.type_embedding = nn.Embedding(2, args.hidden)
layer = nn.TransformerEncoderLayer(
d_model=args.hidden, nhead=args.heads, dim_feedforward=args.feed_forward,
dropout=args.dropout, activation="gelu", batch_first=True, norm_first=True,
)
self.encoder = nn.TransformerEncoder(layer, num_layers=args.layers, enable_nested_tensor=False)
self.norm = nn.LayerNorm(args.hidden)
self.projection = nn.Linear(args.hidden, 768)
def forward(self, input_ids, attention_mask, token_type_ids=None):
if token_type_ids is None:
token_type_ids = torch.zeros_like(input_ids)
positions = torch.arange(input_ids.size(1), device=input_ids.device).unsqueeze(0)
values = (
self.token_embedding(input_ids)
+ self.position_embedding(positions)
+ self.type_embedding(token_type_ids)
)
values = self.encoder(values, src_key_padding_mask=attention_mask == 0)
return self.projection(self.norm(values))
class ExportWrapper(nn.Module):
def __init__(self, model, token_id_map):
super().__init__(); self.model = model
self.register_buffer("token_id_map", token_id_map)
def forward(self, input_ids, attention_mask, token_type_ids):
mapped_input_ids = self.token_id_map[input_ids]
return self.model(mapped_input_ids, attention_mask, token_type_ids)[0]
random.seed(args.seed); np.random.seed(args.seed); torch.manual_seed(args.seed)
if torch.cuda.is_available():
torch.cuda.manual_seed_all(args.seed)
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
metadata = json.loads((args.dataset / "dataset.json").read_text(encoding="utf-8"))
model = Student(metadata["vocab_size"]).to(device)
if args.initial_model:
payload = torch.load(args.initial_model, map_location=device, weights_only=False)
state_dict = payload.get("state_dict") or payload.get("best_state") or payload.get("model")
if state_dict is None:
raise ValueError(f"No model state in {args.initial_model}")
model.load_state_dict(state_dict)
if metadata.get("shards"):
records = list(metadata["shards"])
rng = random.Random(args.seed); rng.shuffle(records)
validation_shards = max(1, round(len(records) * args.validation_fraction)) if len(records) > 1 else 0
validation_records = records[:validation_shards]
training_records = records[validation_shards:] or records
training_data = ShardedDataset(args.dataset, training_records, True, args.seed)
validation_data = ShardedDataset(args.dataset, validation_records, False, args.seed) if validation_records else None
loader = DataLoader(training_data, batch_size=args.batch_size, collate_fn=collate)
validation_loader = DataLoader(validation_data, batch_size=args.batch_size, collate_fn=collate) if validation_data else None
else:
all_data = DistillationDataset(args.dataset)
validation_size = max(1, round(len(all_data) * args.validation_fraction)) if len(all_data) > 1 else 0
indices = list(range(len(all_data))); random.Random(args.seed).shuffle(indices)
validation_data = Subset(all_data, indices[:validation_size]) if validation_size else None
training_data = Subset(all_data, indices[validation_size:] or indices)
loader = DataLoader(training_data, batch_size=args.batch_size, shuffle=True, collate_fn=collate)
validation_loader = DataLoader(validation_data, batch_size=args.batch_size, collate_fn=collate) if validation_data else None
optimizer = torch.optim.AdamW(model.parameters(), lr=args.learning_rate, weight_decay=0.01)
scaler = torch.amp.GradScaler("cuda", enabled=device.type == "cuda")
checkpoint_path = args.output / "checkpoint.latest.pt"
if args.resume:
if not checkpoint_path.exists():
raise FileNotFoundError(f"Cannot resume without {checkpoint_path}")
else:
args.output.mkdir(parents=True, exist_ok=False)
started = time.time(); history = []; best_loss = float("inf"); best_state = None; best_epoch = 0; stale_epochs = 0
start_epoch = 1
if args.resume:
checkpoint = torch.load(checkpoint_path, map_location=device, weights_only=False)
model.load_state_dict(checkpoint["model"])
optimizer.load_state_dict(checkpoint["optimizer"])
scaler.load_state_dict(checkpoint["scaler"])
history = checkpoint["history"]
best_loss = checkpoint["best_loss"]
best_state = checkpoint["best_state"]
best_epoch = checkpoint["best_epoch"]
stale_epochs = checkpoint["stale_epochs"]
start_epoch = checkpoint["epoch"] + 1
def calculate_loss(prediction, target, mask):
active = mask.bool().unsqueeze(-1).expand_as(prediction)
mse = torch.mean((prediction[active] - target[active]) ** 2)
cosine = 1.0 - torch.nn.functional.cosine_similarity(prediction[mask.bool()], target[mask.bool()], dim=-1).mean()
return mse + args.cosine_weight * cosine
for epoch in range(start_epoch, args.epochs + 1):
if isinstance(training_data, ShardedDataset):
training_data.set_epoch(epoch)
model.train(); total = 0.0; batches = 0
for ids, mask, target in loader:
ids, mask, target = ids.to(device), mask.to(device), target.to(device)
optimizer.zero_grad(set_to_none=True)
with torch.amp.autocast("cuda", dtype=torch.float16, enabled=device.type == "cuda"):
prediction = model(ids, mask)
loss = calculate_loss(prediction, target, mask)
scaler.scale(loss).backward(); scaler.unscale_(optimizer)
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
scaler.step(optimizer); scaler.update()
total += float(loss.detach()); batches += 1
if args.max_training_batches and batches >= args.max_training_batches:
break
training_loss = total / batches
validation_loss = training_loss
if validation_loader is not None:
model.eval(); validation_total = 0.0; validation_batches = 0
with torch.no_grad():
for ids, mask, target in validation_loader:
ids, mask, target = ids.to(device), mask.to(device), target.to(device)
with torch.amp.autocast("cuda", dtype=torch.float16, enabled=device.type == "cuda"):
validation_total += float(calculate_loss(model(ids, mask), target, mask))
validation_batches += 1
if args.max_validation_batches and validation_batches >= args.max_validation_batches:
break
validation_loss = validation_total / validation_batches
history.append({"epoch": epoch, "training_loss": training_loss, "validation_loss": validation_loss})
print(f"epoch={epoch} training_loss={training_loss:.8f} validation_loss={validation_loss:.8f}")
if validation_loss < best_loss:
best_loss = validation_loss; best_epoch = epoch; stale_epochs = 0
best_state = {name: value.detach().cpu().clone() for name, value in model.state_dict().items()}
else:
stale_epochs += 1
checkpoint = {
"epoch": epoch, "model": model.state_dict(), "optimizer": optimizer.state_dict(),
"scaler": scaler.state_dict(), "history": history, "best_loss": best_loss,
"best_state": best_state, "best_epoch": best_epoch, "stale_epochs": stale_epochs,
}
temporary_checkpoint = checkpoint_path.with_suffix(".tmp")
torch.save(checkpoint, temporary_checkpoint); temporary_checkpoint.replace(checkpoint_path)
if stale_epochs >= args.early_stopping_patience:
print(f"early_stop epoch={epoch} best_epoch={best_epoch}")
break
if best_state is not None:
model.load_state_dict(best_state)
config = {key: str(value) if isinstance(value, Path) else value for key, value in vars(args).items() if key != "handler"}
torch.save({"state_dict": model.state_dict(), "metadata": metadata, "config": config}, args.output / "student.pt")
model.eval()
export_token_map = torch.from_numpy(np.load(args.dataset / "token_id_map.npy")).long()
wrapper = ExportWrapper(model, export_token_map).cpu().eval()
example_ids = torch.tensor([[2, 3]], dtype=torch.long)
example_mask = torch.ones_like(example_ids); example_types = torch.zeros_like(example_ids)
sequence = torch.export.Dim("sequence", min=2, max=args.max_length)
torch.onnx.export(
wrapper, (example_ids, example_mask, example_types), args.output / "bert.student.fp32.onnx",
input_names=["input_ids", "attention_mask", "token_type_ids"], output_names=["logits"],
dynamic_shapes=({1: sequence}, {1: sequence}, {1: sequence}),
opset_version=18, dynamo=True, external_data=False,
)
report = {
"device": str(device), "parameters": sum(parameter.numel() for parameter in model.parameters()),
"training_samples": len(training_data), "validation_samples": len(validation_data) if validation_data is not None else 0,
"best_epoch": best_epoch, "best_validation_loss": best_loss,
"peak_cuda_bytes": torch.cuda.max_memory_allocated() if device.type == "cuda" else 0,
"seconds": time.time() - started, "onnx_bytes": (args.output / "bert.student.fp32.onnx").stat().st_size,
"history": history, "training_batches_last_epoch": batches,
"validation_batches_last_epoch": validation_batches if validation_loader is not None else 0,
}
(args.output / "training.json").write_text(json.dumps(report, indent=2) + "\n", encoding="utf-8")
print(json.dumps({key: value for key, value in report.items() if key != "history"}, indent=2))
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser()
commands = parser.add_subparsers(dest="command", required=True)
dataset = commands.add_parser("dataset")
dataset.add_argument("--assets", type=Path, required=True)
dataset.add_argument("--corpus", type=Path, required=True)
dataset.add_argument("--output", type=Path, required=True)
dataset.add_argument("--max-vocab", type=int, default=0)
dataset.add_argument("--teacher-batch-size", type=int, default=16)
dataset.add_argument("--teacher-dtype", choices=("float16", "float32"), default="float16")
dataset.add_argument("--shard-size", type=int, default=256)
dataset.add_argument("--reuse-token-map", type=Path)
dataset.set_defaults(handler=build_dataset)
train = commands.add_parser("train")
train.add_argument("--dataset", type=Path, required=True)
train.add_argument("--output", type=Path, required=True)
train.add_argument("--epochs", type=int, default=10)
train.add_argument("--batch-size", type=int, default=8)
train.add_argument("--hidden", type=int, default=256)
train.add_argument("--heads", type=int, default=8)
train.add_argument("--feed-forward", type=int, default=768)
train.add_argument("--layers", type=int, default=4)
train.add_argument("--max-length", type=int, default=256)
train.add_argument("--dropout", type=float, default=0.1)
train.add_argument("--learning-rate", type=float, default=3e-4)
train.add_argument("--cosine-weight", type=float, default=0.1)
train.add_argument("--validation-fraction", type=float, default=0.05)
train.add_argument("--early-stopping-patience", type=int, default=5)
train.add_argument("--max-training-batches", type=int, default=0)
train.add_argument("--max-validation-batches", type=int, default=0)
train.add_argument("--resume", action="store_true")
train.add_argument("--initial-model", type=Path)
train.add_argument("--seed", type=int, default=20260721)
train.set_defaults(handler=train_student)
return parser.parse_args()
def main() -> int:
args = parse_args(); args.handler(args); return 0
if __name__ == "__main__":
sys.exit(main())