436 lines
22 KiB
Python
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())
|