File size: 7,801 Bytes
71c20bd | 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 | """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=<dir containing config.json>.")
ecfg = AutoConfig.from_pretrained(cfg_dir).encoder_config
src = {"encoder_layers": [ParakeetEncoderBlock(ecfg, 0)],
"encode_positions": ParakeetEncoderRelPositionalEncoding(ecfg)}
head = cls(variant, "random", src, blank_bias=0.0)
head.load_state_dict(ck["state_dict"])
head.eval()
return head, ck
def _find_config(start: str | None):
"""`config.json` beside the checkpoint, or one directory up. None if there is none."""
for c in ([start, os.path.dirname(start), os.path.dirname(os.path.dirname(start))]
if start else []):
if c and os.path.exists(os.path.join(c, "config.json")):
return c
return None
def ctc_greedy(logits, in_len, blank=BLANK):
"""Collapse per-frame logits into tokens, with the posterior at each emitting frame.
The standard CTC collapse: drop repeats, then drop blanks. The confidence comes back alongside
because the drafter uses it to decide how far to draft.
"""
import math
logp = F.log_softmax(logits.float(), dim=-1)
conf, best = logp.max(-1)
out = []
for b in range(best.shape[0]):
L = int(in_len[b])
seq, cf, prev = [], [], -1
for t, c in zip(best[b, :L].tolist(), conf[b, :L].tolist()):
if t != prev and t != blank:
seq.append(t)
cf.append(math.exp(c))
prev = t
out.append((seq, cf))
return out
|