"""The CTC draft head: the model definition, and how to load a trained one. Separated from `ctc_train.py` because it is needed in two places that share nothing else. Training builds the head by copying modules out of the base model, which is the only way to get the free initialisation; inference just needs the class and a state dict, and ships inside the published model repository where none of the training machinery exists. Depends on `torch`, `transformers` and `safetensors` and nothing local, so it can be copied into a Hub repo and imported there as-is. """ from __future__ import annotations import os import json import torch import torch.nn as nn import torch.nn.functional as F BLANK = 2 # pad_token_id; verified never to appear in a transcript VOCAB = 16384 ENC_DIM = 1280 DEC_DIM = 1024 class DraftHead(nn.Module): """Encoder frames (B, T, 1280) -> per-frame vocabulary logits (B, T, 16384). Frame-synchronous and non-autoregressive: one forward emits the whole hypothesis, which is the entire point. CTC's conditional-independence assumption makes it a poor transcriber and a fine drafter, because every token it proposes is verified by the real decoder before it can reach the output. """ def __init__(self, variant="adapter", init="free", src=None, blank_bias=5.0): super().__init__() self.variant = variant self.adapters = nn.ModuleList() self.encode_positions = None n_ad = {"adapter": 1, "adapter2": 2}.get(variant, 0) if n_ad: import copy top = src["encoder_layers"][-1] for _ in range(n_ad): self.adapters.append(copy.deepcopy(top)) # A ParakeetEncoderBlock is not self-contained: its attention is relative-position and # reads `position_embeddings` through relative_k_proj, so it must be handed the same # encoding the encoder builds. That encoding is a function of sequence length only, so # a copy of the encoder's own module reproduces it exactly. self.encode_positions = copy.deepcopy(src["encode_positions"]) if variant in ("linear", "linear_free"): self.proj = None self.norm = None self.out = nn.Linear(ENC_DIM, VOCAB) if variant == "linear_free" and init == "free": # Folds the decoder's output map into one matrix, which has to skip decoder.norm. # Measured worse than random init -- kept only so the ablation can be reproduced. Wp, bp = src["proj_w"], src["proj_b"] Wo, bo = src["out_w"], src["out_b"] self.out.weight.data = (Wo @ Wp).contiguous() self.out.bias.data = (Wo @ bp + bo).contiguous() else: self.proj = nn.Linear(ENC_DIM, DEC_DIM) self.norm = nn.LayerNorm(DEC_DIM) self.out = nn.Linear(DEC_DIM, VOCAB) if init == "free": self.proj.weight.data = src["proj_w"].clone() self.proj.bias.data = src["proj_b"].clone() self.norm.weight.data = src["norm_w"].clone() self.norm.bias.data = src["norm_b"].clone() self.out.weight.data = src["out_w"].clone() self.out.bias.data = src["out_b"].clone() with torch.no_grad(): self.out.bias[BLANK] += blank_bias def forward(self, x, mask=None): if self.adapters: pos = self.encode_positions(x) am = None if mask is not None: # the block wants the encoder's 4-D pairwise mask, not the (B, T) frame mask am = mask.unsqueeze(1).expand(-1, x.shape[1], -1) am = (am & am.transpose(1, 2)).unsqueeze(1) for ad in self.adapters: o = ad(x, attention_mask=am, position_embeddings=pos) x = o[0] if isinstance(o, tuple) else o if self.proj is not None: x = self.norm(self.proj(x)) return self.out(x) def n_params(self): return sum(p.numel() for p in self.parameters()) @classmethod def from_checkpoint(cls, path, map_location="cpu", config_dir=None): """Rebuild a trained head without loading the 4.1 GB base checkpoint. Training builds the skeleton by copying modules out of the base model. At inference every one of those tensors is about to be overwritten by the state dict, so the skeleton only needs the right *shapes*, and those come from `config.json` alone. `path` is either the `.pt` a training run writes, or a directory holding `draft_head.safetensors` + `draft_config.json`. The directory form is what gets published: safetensors carries no pickle, so it loads without executing anything. `config_dir` supplies the base model's `config.json`. It defaults to the directory being loaded from when that directory has one -- which is what makes a published repo work on a machine that has never seen the original model. """ from transformers import AutoConfig from transformers.models.parakeet.modeling_parakeet import ( ParakeetEncoderBlock, ParakeetEncoderRelPositionalEncoding) d = None if os.path.isdir(path) or str(path).endswith(".safetensors"): from safetensors.torch import load_file d = path if os.path.isdir(path) else os.path.dirname(path) w = (path if str(path).endswith(".safetensors") else os.path.join(d, "draft_head.safetensors")) ck = json.load(open(os.path.join(d, "draft_config.json"), encoding="utf-8")) ck["state_dict"] = load_file(w, device=map_location) else: ck = torch.load(path, map_location=map_location, weights_only=False) d = os.path.dirname(os.path.abspath(path)) variant = ck["variant"] src = None if variant.startswith("adapter"): cfg_dir = config_dir or _find_config(d) if cfg_dir is None: raise FileNotFoundError( "this head needs the base model's config.json to rebuild its adapter block. " "Pass config_dir=