#!/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())