From 0b7759898c7ec93e22115aafdbe74359212e250b Mon Sep 17 00:00:00 2001 From: Alpamys Date: Thu, 26 Mar 2026 13:49:40 +0500 Subject: [PATCH] =?UTF-8?q?fix:=20address=20review=20findings=20=E2=80=94?= =?UTF-8?q?=20immutable=20rows,=20response=20guard,=20GPU=20cleanup?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - Stop mutating dataset rows in-place in _validate_audio_files (use shallow copy) - Guard _generate_server response parsing against unexpected JSON shape - Add empty dataset guard in _prepare_audio_dataset - Free GPU memory after perplexity scoring in compute_perplexity_scores --- soup_cli/commands/generate.py | 5 ++++- soup_cli/data/loader.py | 3 +-- soup_cli/trainer/sft.py | 6 ++++++ soup_cli/utils/quality.py | 5 +++++ 4 files changed, 16 insertions(+), 3 deletions(-) diff --git a/soup_cli/commands/generate.py b/soup_cli/commands/generate.py index a057553..59e201a 100644 --- a/soup_cli/commands/generate.py +++ b/soup_cli/commands/generate.py @@ -454,7 +454,10 @@ def _generate_server( ) data = response.json() - content = data["choices"][0]["message"]["content"] + try: + content = data["choices"][0]["message"]["content"] + except (KeyError, IndexError, TypeError) as exc: + raise ValueError(f"Unexpected server response format: {exc}") from exc return _parse_json_array(content) diff --git a/soup_cli/data/loader.py b/soup_cli/data/loader.py index 16f7235..7ad8c98 100644 --- a/soup_cli/data/loader.py +++ b/soup_cli/data/loader.py @@ -198,8 +198,7 @@ def _validate_audio_files(data: list[dict], audio_dir: Path) -> list[dict]: if not str(resolved).startswith(str(resolved_base)): traversal += 1 continue - row["audio"] = str(resolved) - valid.append(row) + valid.append({**row, "audio": str(resolved)}) if missing > 0: console.print(f"[yellow]Warning: {missing} rows skipped (missing audio path)[/]") diff --git a/soup_cli/trainer/sft.py b/soup_cli/trainer/sft.py index d283d04..a83bfa0 100644 --- a/soup_cli/trainer/sft.py +++ b/soup_cli/trainer/sft.py @@ -516,6 +516,12 @@ class SFTTrainerWrapper: result["sampling_rate"] = sampling_rate return result + if not dataset["train"]: + raise ValueError( + "Audio training dataset is empty after validation. " + "Check audio file paths and audio_dir." + ) + remove_cols = ["messages", "audio"] train_ds = Dataset.from_list(dataset["train"]).map( load_and_format_audio, diff --git a/soup_cli/utils/quality.py b/soup_cli/utils/quality.py index aafca17..f72d9a8 100644 --- a/soup_cli/utils/quality.py +++ b/soup_cli/utils/quality.py @@ -78,6 +78,11 @@ def compute_perplexity_scores( ppl = math.exp(min(loss_val.item(), 100)) # cap to avoid overflow scores.append(ppl) + # Cleanup GPU memory + del model + if torch.cuda.is_available(): + torch.cuda.empty_cache() + return scores