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