File size: 9,263 Bytes
49525ce
 
 
 
39cac0b
49525ce
 
 
 
 
 
 
39cac0b
49525ce
39cac0b
49525ce
 
 
39cac0b
49525ce
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
39cac0b
 
 
 
 
 
 
 
 
 
 
 
49525ce
 
 
 
39cac0b
 
 
 
 
 
49525ce
 
 
 
39cac0b
 
 
 
 
 
 
49525ce
 
 
39cac0b
 
 
 
 
 
 
49525ce
39cac0b
49525ce
39cac0b
 
 
 
 
 
 
 
49525ce
 
 
39cac0b
 
 
8a96a8e
39cac0b
8a96a8e
 
39cac0b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
49525ce
 
39cac0b
 
49525ce
 
 
39cac0b
49525ce
 
 
 
 
 
 
 
 
 
 
 
39cac0b
49525ce
 
 
 
 
 
 
 
 
 
 
 
 
39cac0b
49525ce
 
 
 
 
 
39cac0b
49525ce
 
 
 
 
 
 
39cac0b
49525ce
 
 
 
 
 
 
39cac0b
49525ce
 
 
 
 
 
 
 
 
 
 
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
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
"""
ml/cli.py - Single entry point for the whole diagnostic pipeline
================================================================
Glues: data build -> synthetic lattice -> stutter model -> pronunciation/articulation ->
fusion/self-calibration -> evaluation into one `ml.cli` command.

Commands:
    download        fetch real corpora from HF Hub
    build-dataset   assemble the unified HF dataset (by-speaker split)
    synth-data      generate high-quality semi-synthetic disfluency lattice dataset
    train           fine-tune wav2vec2 + LoRA stutter classifier with Focal Loss
    eval            out-of-speaker accuracy/precision/recall/F1 + evidence
    fusion-fit      fit and evaluate offline logistic regression vs heuristic fusion
    diagnose        run the full multi-modal diagnosis on an audio file
    self-check      run strict pipeline self-checks (including explicit rabbit->wabbit)

Example:
    python -m ml.cli diagnose --input my_speech.wav --prompt "The weather is nice"
    python -m ml.cli self-check
"""
from __future__ import annotations
import argparse
import json
import sys
from pathlib import Path

import numpy as np


def _load_audio(wav):
    import soundfile as sf
    arr, sr = sf.read(str(wav), dtype="float32")
    if arr.ndim > 1:
        arr = arr.mean(axis=1)
    return arr, sr


# ---------------------------------------------------------------------------
# subcommand handlers
# ---------------------------------------------------------------------------
def cmd_download(args):
    from ml.data.download_corpora import fetch_corpus, CORPUS_LOAD
    keys = args.only or list(CORPUS_LOAD)
    for k in keys:
        rec = fetch_corpus(k, dry_run=args.dry_run)
        print(f"  -> {rec['name']}: {rec.get('num_rows', 'n/a')} rows, "
              f"hf={rec.get('hf_id')}")


def cmd_build_dataset(args):
    from ml.data.make_dataset import build
    build(corpora=args.corpora or None, seed=args.seed)


def cmd_synth_data(args):
    from ml.data.make_synthetic_dataset import load_fluent_source_clips, generate_dataset
    clips = load_fluent_source_clips(args.source)
    merge_path = args.source if args.merge_real else None
    generate_dataset(clips, target_count=args.count, seed=args.seed, save_wavs=not args.no_wavs, merge_real_path=merge_path)


def cmd_train(args):
    from ml.model.stutter_trainer import train
    return train(
        data_dir=args.data,
        out_dir=args.out,
        epochs=args.epochs,
        lr=args.lr,
        batch=args.batch,
        seed=args.seed,
        fp16=not args.no_fp16,
        binary=args.binary,
        balance_train=args.balance,
        focal_gamma=args.focal_gamma,
    )


def cmd_fusion_fit(args):
    from ml.model.fusion_fit import fit_fusion
    fit_fusion(
        data_dir=args.data,
        ckpt_dir=args.ckpt,
        out=args.out,
        device=args.device,
    )


def cmd_eval(args):
    from ml.model.evaluate import evaluate
    evaluate(
        data_dir=args.data,
        ckpt_dir=args.ckpt,
        out=args.out,
        threshold=args.threshold,
        device=args.device,
    )


def cmd_diagnose(args):
    from ml.model.engine import SpeechDiagnosticEngine
    engine = SpeechDiagnosticEngine.get_instance(ckpt_dir=args.ckpt)
    res = engine.diagnose_audio(
        audio_input=args.input,
        target_phrase=args.prompt,
        normal_calibration_audio=args.calibrate,
    )
    if args.json:
        print(json.dumps(res, indent=2, default=str))
    else:
        dec = res["decision"]
        print(f"Overall Classification: {dec['buckets']['overall'].upper()}")
        print(f"Fluency Index:          {dec['fluency_100']} / 100")
        print(f"Confidence Level:       {dec.get('confidence', 'N/A')}")
        print(f"ASR Hypothesis:         \"{res['pronunciation'].get('asr_hypothesis', '')}\"")
        print(f"Pronunciation Score:    {res['pronunciation'].get('pron_score', 0)*100:.1f}%")
        print(f"Total Flaws Detected:   {res['flaws']['total_flaws_count']}")
        print(f"Inference Latency:      {res['latency_ms']} ms")


def cmd_selfcheck(args):
    from ml.model import pron_eval
    print("[1/3] Checking DSP signal conditioning & VAD...")
    test_wave = np.random.randn(16000).astype(np.float32) * 0.1
    cond = pron_eval._filter_dc_rumble(test_wave, 16000)
    assert len(cond) == 16000, "Length mismatch in conditioning"
    norm = pron_eval.normalize_for_neural_inference(cond)
    assert np.max(np.abs(norm)) <= 0.90, "Normalization out of bounds"

    print("[2/3] Checking dynamic programming alignment & explicit rabbit->wabbit rhotacism...")
    align = pron_eval.align_words("the red rabbit", "the wed wabbit")
    # Verify exact word-level substitution status on rabbit vs wabbit
    rabbit_item = next((item for item in align if item["expected"] == "rabbit"), None)
    assert rabbit_item is not None, "Rabbit was omitted from alignment"
    assert rabbit_item["status"] == "substitution", f"rabbit->wabbit was marked {rabbit_item['status']}, expected substitution"
    
    flaws = pron_eval.analyze_speech_flaws("the red rabbit", "the wed wabbit", align, {})
    assert any(e["expected"] == "rabbit" for e in flaws["r_sound_issues"]), "rabbit->wabbit substitution was not captured in r_sound_issues!"
    assert any(e["expected"] == "red" for e in flaws["r_sound_issues"]), "red->wed substitution was not captured in r_sound_issues!"

    print("[3/3] Checking sigmatism detection on sun->thun and sweet->thweet...")
    align_s = pron_eval.align_words("the sweet sun", "the thweet thun")
    flaws_s = pron_eval.analyze_speech_flaws("the sweet sun", "the thweet thun", align_s, {})
    assert any(e["expected"] == "sun" for e in flaws_s["s_sound_issues"]), "sun->thun was not captured in s_sound_issues!"
    assert any(e["expected"] == "sweet" for e in flaws_s["s_sound_issues"]), "sweet->thweet was not captured in s_sound_issues!"

    print("All strict pipeline self-checks PASSED!")


def build_parser() -> argparse.ArgumentParser:
    ap = argparse.ArgumentParser(prog="ml.cli", description="Anvaya Speech Pathology Diagnostic CLI")
    sub = ap.add_subparsers(dest="cmd", required=True)

    p = sub.add_parser("download", help="download external corpora from HuggingFace Hub")
    p.add_argument("--only", nargs="*", help="corpus keys to download")
    p.add_argument("--dry-run", action="store_true")
    p.set_defaults(func=cmd_download)

    p = sub.add_parser("build-dataset", help="assemble unified HF dataset")
    p.add_argument("--corpora", nargs="*")
    p.add_argument("--seed", type=int, default=42)
    p.set_defaults(func=cmd_build_dataset)

    p = sub.add_parser("synth-data", help="generate physical .wav synthetic lattice dataset")
    p.add_argument("--source", default="data/metadata/dataset")
    p.add_argument("--count", type=int, default=4000)
    p.add_argument("--seed", type=int, default=42)
    p.add_argument("--no-wavs", action="store_true")
    p.add_argument("--merge-real", action="store_true")
    p.set_defaults(func=cmd_synth_data)

    p = sub.add_parser("train", help="train stutter LoRA model")
    p.add_argument("--data", default="data/synthetic_lattice/dataset")
    p.add_argument("--out", default="ml/models/stutter")
    p.add_argument("--epochs", type=int, default=5)
    p.add_argument("--lr", type=float, default=3e-5)
    p.add_argument("--batch", type=int, default=8)
    p.add_argument("--seed", type=int, default=42)
    p.add_argument("--no-fp16", action="store_true")
    p.add_argument("--binary", dest="binary", action="store_true", default=True)
    p.add_argument("--no-binary", dest="binary", action="store_false")
    p.add_argument("--balance", type=float, default=0.0)
    p.add_argument("--focal-gamma", type=float, default=2.0)
    p.set_defaults(func=cmd_train)

    p = sub.add_parser("fusion-fit", help="fit and evaluate offline logistic fusion vs heuristic")
    p.add_argument("--data", default="data/synthetic_lattice/dataset")
    p.add_argument("--ckpt", default="ml/models/stutter/stutter_lora")
    p.add_argument("--out", default="reports/ev")
    p.add_argument("--device", default=None)
    p.set_defaults(func=cmd_fusion_fit)

    p = sub.add_parser("eval", help="evaluate out-of-speaker test split")
    p.add_argument("--data", default="data/synthetic_lattice/dataset")
    p.add_argument("--ckpt", default="ml/models/stutter/stutter_lora")
    p.add_argument("--out", default="reports/ev")
    p.add_argument("--threshold", type=float, default=0.5)
    p.add_argument("--device", default=None)
    p.set_defaults(func=cmd_eval)

    p = sub.add_parser("diagnose", help="run multi-modal diagnosis on audio file")
    p.add_argument("--input", required=True)
    p.add_argument("--prompt", default="")
    p.add_argument("--ckpt", default="ml/models/stutter/stutter_lora")
    p.add_argument("--calibrate", default=None, help="path to a 'my normal' wav")
    p.add_argument("--json", action="store_true")
    p.set_defaults(func=cmd_diagnose)

    p = sub.add_parser("self-check", help="run strict self-checks")
    p.set_defaults(func=cmd_selfcheck)
    return ap


def main(argv=None):
    args = build_parser().parse_args(argv)
    return args.func(args)


if __name__ == "__main__":
    main()