Skip to content

Commit a0d8330

Browse files
committed
Enhance reference audio management and model configuration validation
- Introduced a new `reference_audio_storage_root` function to manage storage paths for reference audio files, improving organization and accessibility. - Updated audio model configuration functions to utilize the new reference audio management, ensuring proper handling of voice references and paths. - Enhanced validation logic for audio model configurations, including coercion of numeric strings for `device` and `threads` parameters. - Improved tests to cover new functionalities related to reference audio and validation scenarios, ensuring robust integration with existing workflows.
1 parent a603d32 commit a0d8330

19 files changed

Lines changed: 691 additions & 151 deletions

backend/audio_cpp_runtime.py

Lines changed: 7 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -15,6 +15,7 @@
1515
normalize_default_voice_preset,
1616
normalize_voice_presets,
1717
)
18+
from backend.reference_audio import reference_audio_storage_root
1819
from backend.runtime_env import audio_cpp_library_dirs, build_swap_process_env
1920

2021

@@ -198,15 +199,21 @@ def build_audio_cpp_runtime(
198199
model_row[key] = config[key]
199200
if "model_lazy" in config:
200201
model_row["lazy"] = bool(config["model_lazy"])
202+
reference_root = reference_audio_storage_root(
203+
model_path,
204+
storage_key=model.get("id"),
205+
)
201206
presets = normalize_voice_presets(
202207
config.get("voice_presets"),
203208
model_root=model_path,
209+
reference_root=reference_root,
204210
)
205211
if presets:
206212
model_row["voice_presets"] = presets
207213
default_preset = normalize_default_voice_preset(
208214
config.get("default_voice_preset"),
209215
model_root=model_path,
216+
reference_root=reference_root,
210217
voice_presets=presets,
211218
)
212219
if default_preset is not None:
@@ -240,4 +247,3 @@ def build_audio_cpp_runtime(
240247
"use_model_name": stable_id,
241248
"generic_task_path": f"/upstream/{stable_id}/v1/tasks/run",
242249
}
243-

backend/audio_model_config.py

Lines changed: 54 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -19,9 +19,29 @@
1919
validate_speech_default_references,
2020
validate_voice_presets,
2121
)
22+
from backend.reference_audio import reference_audio_storage_root
2223
from backend.model_config import effective_model_config
2324

2425

26+
_AUDIO_DOCUMENTATION_CONFIG_KEYS = frozenset({"key=value"})
27+
28+
29+
def _is_bogus_audio_config_key(key: Any) -> bool:
30+
name = str(key or "")
31+
return not name or name in _AUDIO_DOCUMENTATION_CONFIG_KEYS or "=" in name
32+
33+
34+
def sanitize_audio_engine_section(section: dict) -> dict:
35+
"""Remove parser/documentation artifacts from a stored audio.cpp engine section."""
36+
if not isinstance(section, dict):
37+
return {}
38+
cleaned = dict(section)
39+
for key in list(cleaned):
40+
if _is_bogus_audio_config_key(key):
41+
cleaned.pop(key, None)
42+
return cleaned
43+
44+
2545
_RESERVED_AUDIO_FLAGS = {
2646
"--config",
2747
"--host",
@@ -158,11 +178,29 @@ def _validate_custom_args(value: Any, errors: List[str]) -> None:
158178
errors.append(f"{flag} is Studio-owned and cannot be set in custom_args")
159179

160180

181+
def _coerce_nonneg_int(value: Any) -> Any:
182+
if isinstance(value, int) and not isinstance(value, bool):
183+
return value
184+
if isinstance(value, float) and not isinstance(value, bool):
185+
return int(value)
186+
if isinstance(value, str):
187+
stripped = value.strip()
188+
if stripped.isdigit() or (
189+
stripped.startswith("-") and stripped[1:].isdigit()
190+
):
191+
return int(stripped)
192+
return value
193+
194+
161195
def _validate_core_runtime_options(config: dict, errors: List[str]) -> None:
162196
for key, minimum in (("device", 0), ("threads", 1)):
163197
value = config.get(key)
164198
if not _present(value):
165199
continue
200+
coerced = _coerce_nonneg_int(value)
201+
if coerced is not value:
202+
config[key] = coerced
203+
value = coerced
166204
if not isinstance(value, int) or isinstance(value, bool):
167205
errors.append(f"{key} must be int")
168206
continue
@@ -188,6 +226,10 @@ def validate_audio_model_config(
188226
Raises ``ValueError`` with all user-actionable validation failures.
189227
"""
190228

229+
engines = normalized_config.get("engines")
230+
if isinstance(engines, dict) and isinstance(engines.get("audio_cpp"), dict):
231+
engines["audio_cpp"] = sanitize_audio_engine_section(engines["audio_cpp"])
232+
191233
effective = effective_model_config(normalized_config)
192234
if effective.get("engine") != "audio_cpp":
193235
return {"errors": [], "warnings": []}
@@ -301,6 +343,15 @@ def validate_audio_model_config(
301343
)
302344

303345
_validate_core_runtime_options(effective, errors)
346+
audio_section = (
347+
normalized_config.get("engines", {}).get("audio_cpp")
348+
if isinstance(normalized_config.get("engines"), dict)
349+
else None
350+
)
351+
if isinstance(audio_section, dict):
352+
for key in ("device", "threads"):
353+
if key in effective:
354+
audio_section[key] = effective[key]
304355

305356
request_options = effective.get("request_options")
306357
if isinstance(request_options, dict) and request_options:
@@ -314,15 +365,18 @@ def validate_audio_model_config(
314365
or model.get("model_path")
315366
or ""
316367
)
368+
reference_root = reference_audio_storage_root(model_root, storage_key=model.get("id"))
317369
if is_tts_task(task):
318370
validate_voice_presets(
319371
effective,
320372
model_root=model_root,
373+
reference_root=reference_root,
321374
errors=errors,
322375
)
323376
validate_speech_default_references(
324377
effective,
325378
model_root=model_root,
379+
reference_root=reference_root,
326380
errors=errors,
327381
)
328382
if effective.get("speech_defaults") is not None and not isinstance(

backend/audio_voice_presets.py

Lines changed: 109 additions & 18 deletions
Original file line numberDiff line numberDiff line change
@@ -16,54 +16,94 @@ def _clean_text(value: Any) -> Optional[str]:
1616
return text or None
1717

1818

19-
def _resolve_voice_ref(model_root: str, value: Any) -> Optional[str]:
19+
def _reference_roots(model_root: str, reference_root: Optional[str] = None) -> list[str]:
20+
roots: list[str] = []
21+
for root in (reference_root, model_root):
22+
text = _clean_text(root)
23+
if text:
24+
resolved = os.path.abspath(text)
25+
if resolved not in roots:
26+
roots.append(resolved)
27+
return roots
28+
29+
30+
def _resolve_voice_ref(
31+
model_root: str,
32+
value: Any,
33+
*,
34+
reference_root: Optional[str] = None,
35+
) -> Optional[str]:
2036
raw = _clean_text(value)
2137
if not raw:
2238
return None
2339
if os.path.isabs(raw):
2440
return raw
25-
root = os.path.abspath(model_root)
41+
roots = _reference_roots(model_root, reference_root)
42+
for root in roots:
43+
candidate = os.path.abspath(os.path.join(root, raw))
44+
if os.path.exists(candidate):
45+
return candidate
46+
root = roots[0] if roots else os.path.abspath(model_root)
2647
return os.path.abspath(os.path.join(root, raw))
2748

2849

29-
def _relative_voice_ref_escapes(model_root: str, value: Any) -> bool:
50+
def _relative_voice_ref_escapes(
51+
model_root: str,
52+
value: Any,
53+
*,
54+
reference_root: Optional[str] = None,
55+
) -> bool:
3056
raw = _clean_text(value)
3157
if not raw or os.path.isabs(raw):
3258
return False
33-
root = os.path.realpath(model_root)
34-
target = os.path.realpath(os.path.join(root, raw))
35-
try:
36-
return os.path.commonpath([root, target]) != root
37-
except ValueError:
38-
return True
59+
escaped_all_roots = True
60+
for root in _reference_roots(model_root, reference_root):
61+
root_real = os.path.realpath(root)
62+
target = os.path.realpath(os.path.join(root_real, raw))
63+
try:
64+
if os.path.commonpath([root_real, target]) == root_real:
65+
escaped_all_roots = False
66+
except ValueError:
67+
continue
68+
return escaped_all_roots
3969

4070

4171
def _validate_voice_ref_path(
4272
*,
4373
label: str,
4474
value: Any,
4575
model_root: str,
76+
reference_root: Optional[str] = None,
4677
errors: list[str],
4778
) -> None:
4879
raw = _clean_text(value)
4980
if not raw:
5081
return
51-
if _relative_voice_ref_escapes(model_root, raw):
82+
if _relative_voice_ref_escapes(model_root, raw, reference_root=reference_root):
5283
errors.append(f"{label} escapes model bundle: {raw}")
5384
return
54-
resolved = _resolve_voice_ref(model_root, raw)
85+
resolved = _resolve_voice_ref(model_root, raw, reference_root=reference_root)
5586
if not resolved or not os.path.exists(resolved):
5687
errors.append(f"{label} does not exist: {resolved or raw}")
5788

5889

59-
def normalize_voice_preset(preset: Any, *, model_root: str) -> Optional[Dict[str, str]]:
90+
def normalize_voice_preset(
91+
preset: Any,
92+
*,
93+
model_root: str,
94+
reference_root: Optional[str] = None,
95+
) -> Optional[Dict[str, str]]:
6096
if not isinstance(preset, dict):
6197
return None
6298
out: Dict[str, str] = {}
6399
voice_id = _clean_text(preset.get("voice_id"))
64100
if voice_id:
65101
out["voice_id"] = voice_id
66-
voice_ref = _resolve_voice_ref(model_root, preset.get("voice_ref"))
102+
voice_ref = _resolve_voice_ref(
103+
model_root,
104+
preset.get("voice_ref"),
105+
reference_root=reference_root,
106+
)
67107
if voice_ref:
68108
out["voice_ref"] = voice_ref
69109
reference_text = _clean_text(preset.get("reference_text"))
@@ -76,6 +116,7 @@ def normalize_voice_presets(
76116
presets: Any,
77117
*,
78118
model_root: str,
119+
reference_root: Optional[str] = None,
79120
) -> Dict[str, Dict[str, str]]:
80121
if not isinstance(presets, dict):
81122
return {}
@@ -84,7 +125,11 @@ def normalize_voice_presets(
84125
name = str(raw_name or "").strip()
85126
if not name:
86127
continue
87-
normalized = normalize_voice_preset(raw_preset, model_root=model_root)
128+
normalized = normalize_voice_preset(
129+
raw_preset,
130+
model_root=model_root,
131+
reference_root=reference_root,
132+
)
88133
if normalized:
89134
out[name] = normalized
90135
return out
@@ -94,6 +139,7 @@ def normalize_default_voice_preset(
94139
value: Any,
95140
*,
96141
model_root: str,
142+
reference_root: Optional[str] = None,
97143
voice_presets: Optional[Dict[str, Dict[str, str]]] = None,
98144
) -> Any:
99145
if value is None or value == "":
@@ -102,7 +148,11 @@ def normalize_default_voice_preset(
102148
name = value.strip()
103149
return name or None
104150
if isinstance(value, dict):
105-
normalized = normalize_voice_preset(value, model_root=model_root)
151+
normalized = normalize_voice_preset(
152+
value,
153+
model_root=model_root,
154+
reference_root=reference_root,
155+
)
106156
return normalized
107157
return None
108158

@@ -173,7 +223,30 @@ def _task_normalized_to_swap_params(normalized: Dict[str, Any]) -> Dict[str, Any
173223
return params
174224

175225

176-
def audio_request_defaults_to_swap_set_params(config: dict) -> Dict[str, Any]:
226+
def _resolve_normalized_reference_fields(
227+
normalized: Dict[str, Any],
228+
*,
229+
model_root: Optional[str] = None,
230+
reference_root: Optional[str] = None,
231+
) -> Dict[str, Any]:
232+
if not model_root or "voice_ref" not in normalized:
233+
return normalized
234+
resolved = _resolve_voice_ref(
235+
model_root,
236+
normalized.get("voice_ref"),
237+
reference_root=reference_root,
238+
)
239+
if not resolved:
240+
return normalized
241+
return {**normalized, "voice_ref": resolved}
242+
243+
244+
def audio_request_defaults_to_swap_set_params(
245+
config: dict,
246+
*,
247+
model_root: Optional[str] = None,
248+
reference_root: Optional[str] = None,
249+
) -> Dict[str, Any]:
177250
"""
178251
Map Studio request defaults to llama-swap ``filters.setParams`` for audio.cpp.
179252
@@ -186,6 +259,11 @@ def audio_request_defaults_to_swap_set_params(config: dict) -> Dict[str, Any]:
186259

187260
defaults_key = request_defaults_key_for(config.get("task"), config.get("family"))
188261
normalized = normalize_request_defaults(defaults_key, config.get(defaults_key))
262+
normalized = _resolve_normalized_reference_fields(
263+
normalized,
264+
model_root=model_root,
265+
reference_root=reference_root,
266+
)
189267
if not normalized:
190268
return {}
191269
if defaults_key == "speech_defaults":
@@ -254,16 +332,25 @@ def validate_voice_presets(
254332
config: dict,
255333
*,
256334
model_root: str,
335+
reference_root: Optional[str] = None,
257336
errors: list[str],
258337
) -> Dict[str, Dict[str, str]]:
259-
presets = normalize_voice_presets(config.get("voice_presets"), model_root=model_root)
338+
presets = normalize_voice_presets(
339+
config.get("voice_presets"),
340+
model_root=model_root,
341+
reference_root=reference_root,
342+
)
260343
default_value = config.get("default_voice_preset")
261344
if isinstance(default_value, str):
262345
name = default_value.strip()
263346
if name and name not in presets:
264347
errors.append(f"default_voice_preset '{name}' is not defined in voice_presets")
265348
elif isinstance(default_value, dict):
266-
if not normalize_voice_preset(default_value, model_root=model_root):
349+
if not normalize_voice_preset(
350+
default_value,
351+
model_root=model_root,
352+
reference_root=reference_root,
353+
):
267354
errors.append(
268355
"default_voice_preset must include voice_id, voice_ref, or reference_text"
269356
)
@@ -273,6 +360,7 @@ def validate_voice_presets(
273360
label=f"default_voice_preset {path_key}",
274361
value=raw,
275362
model_root=model_root,
363+
reference_root=reference_root,
276364
errors=errors,
277365
)
278366
raw_presets = config.get("voice_presets")
@@ -288,6 +376,7 @@ def validate_voice_presets(
288376
label=f"voice_presets.{name}.voice_ref",
289377
value=raw_voice_ref,
290378
model_root=model_root,
379+
reference_root=reference_root,
291380
errors=errors,
292381
)
293382
return presets
@@ -297,6 +386,7 @@ def validate_speech_default_references(
297386
config: dict,
298387
*,
299388
model_root: str,
389+
reference_root: Optional[str] = None,
300390
errors: list[str],
301391
) -> None:
302392
defaults = config.get("speech_defaults")
@@ -306,5 +396,6 @@ def validate_speech_default_references(
306396
label="speech_defaults.voice_ref",
307397
value=defaults.get("voice_ref"),
308398
model_root=model_root,
399+
reference_root=reference_root,
309400
errors=errors,
310401
)

0 commit comments

Comments
 (0)