Spaces:
Configuration error
Configuration error
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() |