File size: 7,067 Bytes
5c2beba
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
176
177
178
179
180
181
182
183
184
185
186
187
#!/usr/bin/env python3
"""Batch-synthesize the Egyptian eval set with s2-pro (+optional LoRA ckpt).

Loads the text2semantic model + codec ONCE, then synthesizes every sentence in
--sentences (jsonl: {id, text}) using a fixed reference voice, writing
<outdir>/<id>.wav plus timing.jsonl with generation stats (RTF etc).

Usage:
  python synth_eval.py --outdir /opt/work/eval/baseline
  python synth_eval.py --outdir /opt/work/eval/step_2000 \
      --lora-ckpt results/s2pro_egy_lora/checkpoints/step_000002000.ckpt \
      --lora-config r_32_egy
"""

import argparse
import json
import time
from pathlib import Path

import numpy as np
import soundfile as sf
import torch
from loguru import logger

from fish_speech.models.text2semantic.inference import (
    encode_audio,
    generate_long,
    init_model,
    load_codec_model,
)

CKPT = Path("/opt/work/checkpoints/s2-pro")


def load_lora(model, ckpt_path: str, lora_config_name: str, lora_filter: str = "all"):
    from hydra import compose, initialize_config_dir
    from hydra.utils import instantiate

    from fish_speech.models.text2semantic.lora import setup_lora

    cfg_dir = str(
        Path("/opt/work/fish-speech/fish_speech/configs/lora").resolve()
    )
    with initialize_config_dir(version_base="1.3", config_dir=cfg_dir):
        lora_cfg = instantiate(compose(config_name=lora_config_name))
    setup_lora(model, lora_cfg)

    sd = torch.load(ckpt_path, map_location="cpu", weights_only=False)
    if "state_dict" in sd:
        sd = sd["state_dict"]
    sd = {k.removeprefix("model."): v for k, v in sd.items()}
    if lora_filter == "slow":
        sd = {k: v for k, v in sd.items() if ".fast_" not in k and not k.startswith("fast_")}
        logger.info(f"ablation slow-only: {len(sd)} tensors kept")
    elif lora_filter == "fast":
        sd = {k: v for k, v in sd.items() if ".fast_" in k or k.startswith("fast_")}
        logger.info(f"ablation fast-only: {len(sd)} tensors kept")
    err = model.load_state_dict(sd, strict=False)
    n_lora = sum(1 for k in sd if "lora" in k)
    assert n_lora > 0, "no lora keys in checkpoint!"
    logger.info(f"Loaded {n_lora} LoRA tensors; missing={len(err.missing_keys)} (expected: base weights)")
    if err.unexpected_keys:
        logger.warning(f"Unexpected keys: {err.unexpected_keys[:5]}")
    return model


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--sentences", default="/opt/work/scripts/eval_sentences.jsonl")
    ap.add_argument("--outdir", required=True)
    ap.add_argument("--lora-ckpt", default=None)
    ap.add_argument("--lora-config", default="r_32_egy")
    ap.add_argument("--ref-audio", default="/opt/work/eval/ref_voice.wav")
    ap.add_argument("--ref-text-file", default="/opt/work/eval/ref_voice.txt")
    ap.add_argument("--temperature", type=float, default=0.8)
    ap.add_argument("--top-p", type=float, default=0.8)
    ap.add_argument("--seed", type=int, default=42)
    ap.add_argument("--skip-existing", action="store_true",
                    help="only synthesize ids whose wav is missing; merge timing")
    ap.add_argument("--lora-filter", choices=["all", "slow", "fast"], default="all",
                    help="ablation: load only slow-transformer or only fast-transformer LoRA weights")
    ap.add_argument("--limit", type=int, default=0, help="synth only first N sentences")
    args = ap.parse_args()

    outdir = Path(args.outdir)
    outdir.mkdir(parents=True, exist_ok=True)
    device = "cuda"
    precision = torch.bfloat16

    sentences = [
        json.loads(l)
        for l in Path(args.sentences).read_text(encoding="utf-8").splitlines()
        if l.strip()
    ]
    old_timing = {}
    tj = outdir / "timing.jsonl"
    if args.skip_existing and tj.exists():
        for l in tj.read_text(encoding="utf-8").splitlines():
            if l.strip():
                r = json.loads(l)
                old_timing[r["id"]] = r
        before = len(sentences)
        sentences = [s for s in sentences
                     if not (outdir / f"{s['id']}.wav").exists()]
        logger.info(f"skip-existing: {before - len(sentences)} kept, {len(sentences)} to synth")
        if not sentences:
            logger.info("nothing to do")
            return
    if args.limit:
        sentences = sentences[: args.limit]

    logger.info("Loading text2semantic model...")
    model, decode_one_token = init_model(CKPT, device, precision, compile=False)
    if args.lora_ckpt:
        model = load_lora(model, args.lora_ckpt, args.lora_config, args.lora_filter)
        model = model.to(device=device, dtype=precision).eval()
    with torch.device(device):
        model.setup_caches(
            max_batch_size=1,
            max_seq_len=model.config.max_seq_len,
            dtype=next(model.parameters()).dtype,
        )

    logger.info("Loading codec...")
    codec = load_codec_model(CKPT / "codec.pth", device, precision)

    ref_text = Path(args.ref_text_file).read_text(encoding="utf-8").strip()
    ref_tokens = encode_audio(args.ref_audio, codec, device).cpu()

    results = []
    for s in sentences:
        torch.manual_seed(args.seed)
        torch.cuda.manual_seed(args.seed)
        t0 = time.time()
        gen = generate_long(
            model=model,
            device=device,
            decode_one_token=decode_one_token,
            text=s["text"],
            num_samples=1,
            max_new_tokens=0,
            top_p=args.top_p,
            top_k=30,
            temperature=args.temperature,
            compile=False,
            iterative_prompt=True,
            chunk_length=300,
            prompt_text=[ref_text],
            prompt_tokens=[ref_tokens],
        )
        codes = []
        for r in gen:
            if r.action == "sample":
                codes.append(r.codes)
        gen_s = time.time() - t0
        if not codes:
            logger.error(f"{s['id']}: NO CODES GENERATED")
            results.append({"id": s["id"], "error": "no_codes"})
            continue
        merged = torch.cat(codes, dim=1).to(device)
        with torch.no_grad():
            fake = codec.from_indices(merged.unsqueeze(0))
        wav = fake[0, 0].float().cpu().numpy()
        dur = len(wav) / codec.sample_rate
        sf.write(outdir / f"{s['id']}.wav", wav, codec.sample_rate)
        rtf = gen_s / max(dur, 1e-6)
        results.append(
            {"id": s["id"], "text": s["text"], "gen_s": round(gen_s, 2),
             "dur_s": round(dur, 2), "rtf": round(rtf, 2)}
        )
        logger.info(f"{s['id']}: {dur:.1f}s audio in {gen_s:.1f}s (RTF {rtf:.2f})")

    merged = {**old_timing, **{r["id"]: r for r in results}}
    with open(outdir / "timing.jsonl", "w", encoding="utf-8") as f:
        for r in merged.values():
            f.write(json.dumps(r, ensure_ascii=False) + "\n")
    ok = [r for r in results if "rtf" in r]
    if ok:
        logger.info(
            f"DONE {len(ok)}/{len(results)} ok; mean RTF "
            f"{sum(r['rtf'] for r in ok)/len(ok):.2f}"
        )


if __name__ == "__main__":
    main()