Skip to content

Commit 4ce5a1c

Browse files
authored
Merge pull request #48 from Subohao/metric_description
Adding Text-LLMs description functionalities to VERSA
2 parents 3239d7c + 6a734f3 commit 4ce5a1c

7 files changed

Lines changed: 869 additions & 0 deletions

File tree

‎egs/interpreter.yaml‎

Lines changed: 8 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,8 @@
1+
# interpreter example yaml config
2+
# A list of interpreter backends your code will load.
3+
# Each item must have at least `model_name`.
4+
# For some models (Mistral, Llama 3.1), you must also provide HF_TOKEN.
5+
6+
interpreter_config:
7+
# Easiest path: no HF login required in your loader
8+
- model_name: "Qwen/Qwen2.5-7B-Instruct"

‎scripts/chunk_func/chunk.py‎

Lines changed: 145 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,145 @@
1+
#!/usr/bin/env python3
2+
3+
# Copyright 2025 BoHao Su
4+
# Apache 2.0 (http://www.apache.org/licenses/LICENSE-2.0)
5+
6+
import argparse
7+
import os
8+
from pathlib import Path
9+
10+
import numpy as np
11+
import soundfile as sf
12+
from tqdm import tqdm
13+
from versa.scorer_shared import audio_loader_setup, load_audio, wav_normalize
14+
15+
16+
def get_parser() -> argparse.Namespace:
17+
"""Get argument parser."""
18+
parser = argparse.ArgumentParser(description="Chunk audios into fixed durations.")
19+
parser.add_argument(
20+
"--pred",
21+
type=str,
22+
required=True,
23+
help="Wav.scp for generated waveforms, or a dir depending on --io.",
24+
)
25+
parser.add_argument(
26+
"--io",
27+
type=str,
28+
default="kaldi",
29+
choices=["kaldi", "soundfile", "dir"],
30+
help="IO interface to use.",
31+
)
32+
parser.add_argument(
33+
"--chunk_duration",
34+
type=float,
35+
default=3.0,
36+
help="Duration (sec) of each chunk window.",
37+
)
38+
parser.add_argument(
39+
"--hop_duration",
40+
type=float,
41+
default=None,
42+
help="Hop size (sec) between chunk starts. "
43+
"If None, equals --chunk_duration (non-overlap).",
44+
)
45+
parser.add_argument(
46+
"--output_dir",
47+
type=str,
48+
required=True,
49+
help="Directory to write chunked wav files.",
50+
)
51+
parser.add_argument(
52+
"--min_last_chunk",
53+
type=float,
54+
default=0.0,
55+
help="Minimum duration (sec) required to keep the final (short) chunk. "
56+
"Set >0 to drop very short tails.",
57+
)
58+
return parser
59+
60+
61+
def main():
62+
args = get_parser().parse_args()
63+
64+
output_dir = Path(args.output_dir)
65+
output_dir.mkdir(parents=True, exist_ok=True)
66+
67+
if args.chunk_duration <= 0:
68+
raise ValueError("--chunk_duration must be > 0")
69+
70+
hop_duration = (
71+
args.hop_duration if args.hop_duration is not None else args.chunk_duration
72+
)
73+
if hop_duration <= 0:
74+
raise ValueError("--hop_duration must be > 0")
75+
76+
if args.min_last_chunk < 0:
77+
raise ValueError("--min_last_chunk must be >= 0")
78+
79+
gen_files = audio_loader_setup(args.pred, args.io)
80+
if len(gen_files) == 0:
81+
raise FileNotFoundError(
82+
"Not found any generated audio files from --pred with --io."
83+
)
84+
85+
total_chunks = 0
86+
for key in tqdm(list(gen_files.keys()), desc="Chunking"):
87+
src_path = gen_files[key]
88+
try:
89+
sr, wav = load_audio(src_path, args.io)
90+
wav = wav_normalize(wav)
91+
if wav.ndim > 1:
92+
# Convert to mono if multichannel
93+
wav = np.mean(wav, axis=-1)
94+
except Exception as e:
95+
print(f"[WARN] Failed to load {key} from {src_path}: {e}")
96+
continue
97+
98+
chunk_len = int(round(args.chunk_duration * sr))
99+
hop_len = int(round(hop_duration * sr))
100+
min_last_len = int(round(args.min_last_chunk * sr))
101+
102+
if chunk_len <= 0 or hop_len <= 0:
103+
print(f"[WARN] Non-positive chunk/hop for key={key}; skipping.")
104+
continue
105+
106+
n_samples = len(wav)
107+
if n_samples == 0:
108+
print(f"[WARN] Empty audio for key={key}; skipping.")
109+
continue
110+
111+
# Iterate chunk start positions
112+
chunk_idx = 0
113+
start = 0
114+
while start < n_samples:
115+
end = start + chunk_len
116+
if end > n_samples:
117+
# last (short) chunk
118+
if (n_samples - start) < min_last_len:
119+
break # drop the tail if too short
120+
end = n_samples
121+
122+
chunk = wav[start:end]
123+
if len(chunk) == 0:
124+
break
125+
126+
# Include time range in filename for traceability
127+
t0 = start / sr
128+
t1 = end / sr
129+
out_name = f"{key}_chunk{chunk_idx:04d}_{t0:.3f}-{t1:.3f}.wav"
130+
out_path = output_dir / out_name
131+
132+
try:
133+
sf.write(str(out_path), chunk, sr, subtype="PCM_16")
134+
total_chunks += 1
135+
except Exception as e:
136+
print(f"[WARN] Failed to write {out_path}: {e}")
137+
138+
chunk_idx += 1
139+
start += hop_len
140+
141+
print(f"Done. Wrote {total_chunks} chunks to: {output_dir.resolve()}")
142+
143+
144+
if __name__ == "__main__":
145+
main()

‎scripts/description/README.md‎

Lines changed: 36 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,36 @@
1+
# Speech Evaluation Interpreter
2+
3+
This tool loads utterance-level metrics and uses LLM interpreters to generate natural-language descriptions.
4+
5+
6+
## Files
7+
- `interpreter.py`: CLI entry point, loads config, metrics, runs interpreters, saves JSON.
8+
- `interpreter_shared.py`: utilities for loading metrics and models.
9+
- `text_llm_description.py`: **you implement** `describe_all(...)` to describe each utterance.
10+
11+
12+
## Example Input
13+
14+
### scores.scp
15+
16+
```
17+
utt_0001 {"SNR": 23.1, "WER": 0.08, "MOS": 4.2}
18+
utt_0002 {"SNR": 12.7, "WER": 0.30, "MOS": 3.0}
19+
```
20+
21+
### egs/interpreter.yaml
22+
```yaml
23+
interpreter_config:
24+
- model_name: "Qwen/Qwen2.5-7B-Instruct"
25+
```
26+
27+
## Run
28+
29+
```bash
30+
python interpreter.py \
31+
--config egs/interpreter.yaml \
32+
--score_output_file scores.scp \
33+
--output_file descriptions.json \
34+
--use_gpu False \
35+
--verbose 1
36+
```

‎scripts/description/interpreter.py‎

Lines changed: 108 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,108 @@
1+
#!/usr/bin/env python3
2+
3+
# Copyright 2025 BoHao Su
4+
# Apache 2.0 (http://www.apache.org/licenses/LICENSE-2.0)
5+
6+
"""Interpreter Interface for Speech Evaluation."""
7+
8+
import argparse
9+
import json
10+
import logging
11+
12+
import torch
13+
import yaml
14+
from text_llm_description import describe_all
15+
from interpreter_shared import load_interpreter_modules, metric_loader_setup
16+
17+
18+
def get_parser() -> argparse.Namespace:
19+
"""Get argument parser."""
20+
parser = argparse.ArgumentParser(
21+
description="Interpretation for Speech Evaluation Interface"
22+
)
23+
parser.add_argument(
24+
"--score_output_file",
25+
type=str,
26+
default=None,
27+
help="Path of directory of the score results.",
28+
)
29+
parser.add_argument(
30+
"--config",
31+
required=True,
32+
help="YAML with interpreter_config (list of model_name dicts)",
33+
)
34+
parser.add_argument(
35+
"--output_file", required=True, help="Where to dump the JSON descriptions"
36+
)
37+
parser.add_argument(
38+
"--use_gpu", type=bool, default=False, help="whether to use GPU if it can"
39+
)
40+
parser.add_argument(
41+
"--verbose",
42+
default=1,
43+
type=int,
44+
help="Verbosity level. Higher is more logging.",
45+
)
46+
parser.add_argument(
47+
"--rank",
48+
default=0,
49+
type=int,
50+
help="the overall rank in the batch processing, used to specify GPU rank",
51+
)
52+
return parser
53+
54+
55+
def main():
56+
args = get_parser().parse_args()
57+
58+
# In case of using `local` backend, all GPU will be visible to all process.
59+
if args.use_gpu:
60+
gpu_rank = args.rank % torch.cuda.device_count()
61+
torch.cuda.set_device(gpu_rank)
62+
logging.info(f"using device: cuda:{gpu_rank}")
63+
64+
# logging info
65+
if args.verbose > 1:
66+
logging.basicConfig(
67+
level=logging.DEBUG,
68+
format="%(asctime)s (%(module)s:%(lineno)d) %(levelname)s: %(message)s",
69+
)
70+
elif args.verbose > 0:
71+
logging.basicConfig(
72+
level=logging.INFO,
73+
format="%(asctime)s (%(module)s:%(lineno)d) %(levelname)s: %(message)s",
74+
)
75+
else:
76+
logging.basicConfig(
77+
level=logging.WARN,
78+
format="%(asctime)s (%(module)s:%(lineno)d) %(levelname)s: %(message)s",
79+
)
80+
logging.warning("Skip DEBUG/INFO messages")
81+
82+
metrics = metric_loader_setup(args.score_output_file)
83+
logging.info("The number of utterances = %d" % len(metrics))
84+
85+
# 2) Load interpreter modules from YAML
86+
with open(args.config) as cf:
87+
cfg = yaml.safe_load(cf)
88+
interpreter_modules = load_interpreter_modules(
89+
cfg["interpreter_config"],
90+
use_gpu=args.use_gpu,
91+
)
92+
93+
# 3) Run description for each model
94+
all_results = []
95+
for model_cfg in cfg["interpreter_config"]:
96+
name = model_cfg["model_name"]
97+
logging.info(f"Describing with {name}")
98+
res = describe_all(metrics, name, interpreter_modules)
99+
all_results.extend(res)
100+
101+
# 4) Dump
102+
with open(args.output_file, "w") as outf:
103+
json.dump(all_results, outf, ensure_ascii=False, indent=2)
104+
logging.info(f"Wrote descriptions to {args.output_file}")
105+
106+
107+
if __name__ == "__main__":
108+
main()
Lines changed: 77 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,77 @@
1+
#!/usr/bin/env python3
2+
3+
# Copyright 2025 BoHao Su
4+
# Apache 2.0 (http://www.apache.org/licenses/LICENSE-2.0)
5+
6+
import json
7+
8+
import torch
9+
from huggingface_hub import login
10+
from transformers import AutoModelForCausalLM, AutoTokenizer, pipeline
11+
12+
13+
def metric_loader_setup(score_output_file):
14+
"""
15+
Reads an scp-like file where each line is:
16+
utt_id <JSON-metrics>
17+
Returns a dict mapping utt_id → metrics dict.
18+
"""
19+
data = {}
20+
with open(score_output_file, "r") as f:
21+
for line in f:
22+
line = line.strip()
23+
if not line:
24+
continue
25+
utt_id, json_str = line.split(maxsplit=1)
26+
data[utt_id] = json.loads(json_str)
27+
return data
28+
29+
30+
def load_interpreter_modules(interpreter_config, use_gpu):
31+
assert interpreter_config, "no interpreter function is provided"
32+
interpreter_modules = {}
33+
for config in interpreter_config:
34+
print(config, flush=True)
35+
if config["model_name"] == "Qwen/Qwen2.5-7B-Instruct":
36+
model = AutoModelForCausalLM.from_pretrained(
37+
config["model_name"],
38+
torch_dtype="auto",
39+
device_map="auto" if use_gpu else None,
40+
)
41+
tokenizer = AutoTokenizer.from_pretrained(config["model_name"])
42+
interpreter_modules[config["model_name"]] = {
43+
"args": {
44+
"model": model,
45+
"tokenizer": tokenizer,
46+
},
47+
}
48+
elif config["model_name"] == "mistralai/Mistral-7B-Instruct-v0.3":
49+
login(token=config["HF_TOKEN"])
50+
model = AutoModelForCausalLM.from_pretrained(
51+
config["model_name"],
52+
torch_dtype="auto",
53+
device_map="auto" if use_gpu else None,
54+
)
55+
tokenizer = AutoTokenizer.from_pretrained(config["model_name"])
56+
interpreter_modules[config["model_name"]] = {
57+
"args": {
58+
"model": model,
59+
"tokenizer": tokenizer,
60+
},
61+
}
62+
elif config["model_name"] == "meta-llama/Llama-3.1-8B-Instruct":
63+
login(token=config["HF_TOKEN"])
64+
pipe = pipeline(
65+
"text-generation",
66+
model=config["model_name"],
67+
torch_dtype=torch.bfloat16,
68+
device_map="auto" if use_gpu else None,
69+
)
70+
interpreter_modules[config["model_name"]] = {
71+
"args": {
72+
"pipe": pipe,
73+
},
74+
}
75+
else:
76+
raise ValueError(f"Unsupported model_name: {config['model_name']}")
77+
return interpreter_modules

0 commit comments

Comments
 (0)