Download draft_head.py from sayedM/cohere-transcribe-arabic-cpu-friendly: direct link, hf CLI and curl.
- Browser
- Download file 7.8 kB
-
https://huggingface.co/sayedM/cohere-transcribe-arabic-cpu-friendly/resolve/main/draft_head.py
- Command line
-
hf download hf://sayedM/cohere-transcribe-arabic-cpu-friendly/draft_head.py
-
curl -L -o draft_head.py https://huggingface.co/sayedM/cohere-transcribe-arabic-cpu-friendly/resolve/main/draft_head.py
7.8 kB
| """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()) | |
| 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 | |