File size: 6,882 Bytes
f9e3832
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
174
175
"""Stage 1: precompute (audio_embeds, teacher_token_ids) per clip on CPU.

For each audio clip in the training corpus:
  1. Load audio, run Voxtral audio_tower + multi_modal_projector to get audio_embeds.
  2. Run full Voxtral end-to-end on CPU (greedy) to get teacher's transcript
     token IDs. This is the "pseudo-label" for distillation.
  3. Also save the full input_ids sequence and the audio_token_id mask so
     downstream code can splice audio_embeds into the LLM input.
  4. Save each clip as `<clip>.pt`.

CPU-only. This is slow (Voxtral is 3B params on CPU) but only runs once per
clip and can be run offline in parallel while we develop the training loop.
Only rerun if the target model or prompt template changes.

Usage:
    python precompute_pseudo_labels.py \
        --model mistralai/Voxtral-Mini-3B-2507 \
        --manifest /mnt/data/librispeech_manifest.csv \
        --dataset-root / \
        --output-dir /mnt/data/pseudo_labels \
        --language en \
        --max-new-tokens 256

Output per clip: /mnt/data/pseudo_labels/<clip_name>.pt with keys:
    - input_ids: [1, prefix_len] LongTensor (Voxtral prompt + audio placeholder tokens)
    - audio_embeds: [n_audio_tokens, hidden_size] BFloat16Tensor
    - audio_token_id: int
    - teacher_token_ids: [n_generated] LongTensor (target's greedy transcript)
    - reference: str (ground-truth transcript, informational only)
    - duration_sec: float
    - metadata: {model_id, language, prompt_variant}
"""

from __future__ import annotations

import argparse
import csv
import gc
import json
import time
from pathlib import Path

import torch
from transformers import VoxtralForConditionalGeneration, AutoProcessor


def parse_args() -> argparse.Namespace:
    p = argparse.ArgumentParser()
    p.add_argument("--model", default="mistralai/Voxtral-Mini-3B-2507")
    p.add_argument("--manifest", required=True, type=Path)
    p.add_argument("--dataset-root", required=True, type=Path)
    p.add_argument("--output-dir", required=True, type=Path)
    p.add_argument("--language", default="en")
    p.add_argument("--max-new-tokens", type=int, default=256)
    p.add_argument("--limit", type=int, default=None, help="Limit to first N clips")
    p.add_argument("--skip-existing", action="store_true", default=True)
    return p.parse_args()


def read_manifest(path: Path, root: Path, limit: int | None):
    rows = []
    with open(path) as f:
        r = csv.DictReader(f)
        for row in r:
            rel = row["audio_path"]
            rows.append({
                "audio_path": root / rel,
                "clip_name": Path(rel).stem,
                "duration_sec": float(row["duration_sec"]),
                "reference": row["transcript"],
            })
            if limit is not None and len(rows) >= limit:
                break
    return rows


def main() -> int:
    args = parse_args()
    torch.set_grad_enabled(False)

    print(f"[precompute] Loading {args.model} on CPU (bfloat16)")
    dtype = torch.bfloat16
    t0 = time.perf_counter()
    model = VoxtralForConditionalGeneration.from_pretrained(
        args.model, torch_dtype=dtype, low_cpu_mem_usage=True,
    ).eval()
    processor = AutoProcessor.from_pretrained(args.model)
    print(f"[precompute] Loaded in {time.perf_counter()-t0:.1f}s")

    audio_token_id = model.config.audio_token_id
    rows = read_manifest(args.manifest, args.dataset_root, args.limit)
    print(f"[precompute] {len(rows)} clips to process")

    args.output_dir.mkdir(parents=True, exist_ok=True)

    metadata = {
        "model_id": args.model,
        "language": args.language,
        "max_new_tokens": args.max_new_tokens,
        "audio_token_id": audio_token_id,
        "hidden_size": model.config.text_config.hidden_size,
        "vocab_size": model.config.text_config.vocab_size,
        "num_hidden_layers": model.config.text_config.num_hidden_layers,
    }
    with open(args.output_dir / "_metadata.json", "w") as f:
        json.dump(metadata, f, indent=2)

    processed = 0
    for i, row in enumerate(rows):
        out_path = args.output_dir / f"{row['clip_name']}.pt"
        if args.skip_existing and out_path.exists():
            print(f"[precompute] {i+1}/{len(rows)}: {row['clip_name']}.pt exists, skipping")
            processed += 1
            continue

        print(f"[precompute] {i+1}/{len(rows)}: {row['clip_name']} ({row['duration_sec']:.2f}s)")
        t_clip = time.perf_counter()

        # -- Build transcription prompt
        inputs = processor.apply_transcription_request(
            language=args.language,
            audio=str(row["audio_path"]),
            model_id=args.model,
        )
        input_ids = inputs["input_ids"]  # [1, prefix_len]
        input_features = inputs["input_features"].to(dtype)  # [1, 128, 3000]

        # -- Compute audio embeds
        with torch.no_grad():
            audio_out = model.model.get_audio_features(input_features)
            audio_embeds = audio_out.pooler_output  # [N_audio_tokens, hidden]

        # -- Run greedy generation to get teacher_token_ids
        t_gen = time.perf_counter()
        with torch.no_grad():
            out = model.generate(
                input_ids=input_ids,
                input_features=input_features,
                max_new_tokens=args.max_new_tokens,
                do_sample=False,
                temperature=None,
                top_p=None,
                pad_token_id=processor.tokenizer.pad_token_id if processor.tokenizer.pad_token_id else processor.tokenizer.eos_token_id,
            )
        teacher_token_ids = out[0, input_ids.shape[1]:].clone()  # [n_generated]
        gen_t = time.perf_counter() - t_gen

        teacher_text = processor.tokenizer.decode(teacher_token_ids, skip_special_tokens=True).strip()
        print(f"[precompute]   audio_embeds={tuple(audio_embeds.shape)}, "
              f"prefix_len={input_ids.shape[1]}, "
              f"teacher_tokens={teacher_token_ids.shape[0]}, "
              f"gen_time={gen_t:.1f}s")
        print(f"[precompute]   teacher: {teacher_text[:100]!r}")

        # -- Save
        torch.save({
            "input_ids": input_ids.cpu(),  # [1, prefix_len]
            "audio_embeds": audio_embeds.cpu().to(dtype),  # [n_audio, hidden]
            "audio_token_id": audio_token_id,
            "teacher_token_ids": teacher_token_ids.cpu(),  # [n_generated]
            "teacher_text": teacher_text,
            "reference": row["reference"],
            "clip_name": row["clip_name"],
            "duration_sec": row["duration_sec"],
        }, out_path)
        processed += 1
        print(f"[precompute]   saved {out_path} (clip time: {time.perf_counter()-t_clip:.1f}s)")

    print(f"\n[precompute] Done. {processed}/{len(rows)} clips written to {args.output_dir}")
    return 0


if __name__ == "__main__":
    raise SystemExit(main())