#!/usr/bin/env python3 """Run ACE-Step for one prepared pack item — the proven reference implementation. This file is the ONLY place GenerationParams is constructed. That is not a style preference, it is the lesson from 2026-07-16/17: apps/stack/infer_worker.py was a second, hand-typed copy of this call, it drifted in six places (global_caption='', audio_cover_strength=0.0, hardcoded bpm, wrong captions, timesignature='', steps=16), and it cost Francesco a night of "still crappy" output. Two hand-typed copies of the same call drift forever and you can never prove you have found the last divergence. One implementation is falsifiable; two are not. Three entry points, one code path: init_ace() load the model ONCE. Expensive (~20-40s: 4B decoder). build_params() metadata.json -> GenerationParams. THE single construction. run_take() one generation against an already-loaded session. ~7s warm. main() the CLI, unchanged. scripts/stemgen_webapp.py and apps/stems shell out to it and must keep working — it is the reference for docs/INFERENCE_RECIPE.md. The CLI does init_ace() + run_take() and exits, so it pays the load every time. A resident caller (apps/stack/infer_worker.py) does init_ace() once and then run_take() per request — the same functions, so it cannot drift from the CLI by construction. """ from __future__ import annotations import argparse import json import os import shutil import sys import time from dataclasses import dataclass from pathlib import Path from typing import Any # ========================================================================= # THE MAY 7 RECIPE — the defaults, living in the ONE implementation. # # Provenance: artifacts/generated_batches/batch_10_actual_lego_20260507/*/ # ace_lego_corrected/result.json — the last batch Francesco confirmed sounds good. # Copied, not invented. docs/INFERENCE_RECIPE.md is the writeup. # # WHY THESE LIVE HERE AND NOT IN A CALLER (found 2026-07-17, live): # apps/stack/infer_worker.py held these as its own constants. When serving switched # to shelling out to this script, the worker stopped being the path — and the recipe # went with it, because this script reads caption/global_caption from metadata.json # and the server's metadata.json has no `global_caption` key. So serving silently # reverted to global_caption="" — the single biggest prompt-side defect — with the # recipe still sitting "restored" in a file nothing called. # # A default in a caller is a default that gets lost. These are the defaults now, so # every caller (CLI, studio, resident worker) gets the proven recipe unless its # metadata deliberately overrides it. # ========================================================================= # The instruction that makes the model emit a STEM rather than a mix. RECIPE_GLOBAL_CAPTION = ( "Generate only the requested missing isolated stem so that it fits the provided " "audio context. Preserve timing, style, tempo, harmony, and arrangement. Do not " "generate a full mix." ) # Captions in ACE-Step's own register: short, natural descriptions of the target # stem — the way ACE was trained, not verbose instructions with negation lists. The # "isolated stem, no full mix" intent is carried by RECIPE_GLOBAL_CAPTION, so the # per-role caption just names the instrument and how it sits. (Was: long prose with # "Do not generate bass/guitars/vocals…", out of ACE's training distribution — # Francesco 2026-07-17. The proven long form is preserved in git if we A/B back.) RECIPE_CAPTIONS = { "drums": "a tight drum kit locked to the groove and tempo", "bass": "a groovy bass line locked to the drums and harmony", "melody": "a melodic lead line locked to the harmony", "vocals": "a lead vocal locked to the melody and phrasing", } RECIPE_TIMESIGNATURE = "4" def _meta_float(meta: dict, key: str): value = meta.get(key) if value in (None, "", "N/A"): return None try: return float(value) except Exception: return None def _meta_int(meta: dict, key: str): value = _meta_float(meta, key) return int(round(value)) if value else None @dataclass class AceSession: """A loaded ACE model. Hold one of these and run_take() is ~7s instead of ~45s.""" dit_handler: Any llm_handler: Any device: str lora_path: str | None = None lora_scale: float = 1.0 full_ft_checkpoint: str | None = None load_seconds: float = 0.0 def init_ace( *, ace_root: str = "/home/fcolo/ace-step-1.5-xl", checkpoints: str = "/home/fcolo/ace-step/checkpoints", model: str = "acestep-v15-xl-base", device: str = "cuda", lm_model: str = "acestep-5Hz-lm-1.7B", lm_backend: str = "pt", no_thinking: bool = True, use_lm: bool = False, use_cot: bool = False, full_ft_checkpoint: str | None = None, lora_path: str | None = None, adapter_name: str = "stemgen", lora_scale: float = 1.0, log=print, ) -> AceSession: """Everything expensive, once. Safe to call from a long-lived process.""" if full_ft_checkpoint and lora_path: raise ValueError("--full-ft-checkpoint and --lora-path are mutually exclusive") t0 = time.time() # Hard CPU mode: ACE/PEFT sometimes tries to stage LoRA weights on CUDA even when # the requested generation device is CPU. Hide CUDA before importing ACE/torch so # CPU jobs don't fight long-running GPU work. if str(device).lower().split(":", 1)[0] == "cpu": os.environ["CUDA_VISIBLE_DEVICES"] = "" os.environ.setdefault("ACESTEP_VAE_ON_CPU", "1") os.environ.setdefault("PYTORCH_ENABLE_MPS_FALLBACK", "1") log("[INFO] CPU mode: CUDA_VISIBLE_DEVICES cleared for this process") ace_root_p = Path(ace_root).resolve() if str(ace_root_p) not in sys.path: sys.path.insert(0, str(ace_root_p)) os.environ["ACESTEP_CHECKPOINTS_DIR"] = checkpoints from acestep.handler import AceStepHandler from acestep.llm_inference import LLMHandler dit_handler = AceStepHandler() status, success = dit_handler.initialize_service( project_root=str(ace_root_p), config_path=model, device=device, prefer_source="huggingface", ) if not success: raise RuntimeError(f"ACE init failed: {status}") log(status) if full_ft_checkpoint: # Full fine-tune: swap the whole DiT decoder. Traced 2026-07-17 -- # initialize_service() -> init_service_loader.py:175 sets # self.model = AutoModel.from_pretrained(...), the SAME construction as # training_v2/model_loader.py:load_decoder_for_training(); and # AceStepConditionGenerationModel.__init__ (modeling_acestep_v15_base.py:1609) # sets self.decoder = AceStepDiTModel(config). So handler.model.decoder is # exactly the submodule train_full.py checkpoints. from safetensors.torch import load_file ckpt = Path(full_ft_checkpoint) / "decoder_model.safetensors" if not ckpt.is_file(): raise RuntimeError(f"no decoder_model.safetensors under {full_ft_checkpoint}") target = getattr(getattr(dit_handler, "model", None), "decoder", None) if target is None: raise RuntimeError("handler.model.decoder not found after initialize_service — " "ACE internals changed; re-trace before trusting this path.") # Fingerprint before/after: a load that silently no-ops (wrong keys, empty # dict) is otherwise indistinguishable from success. probe = next(k for k, _ in target.named_parameters()) before = float(dict(target.named_parameters())[probe].detach().float().sum().item()) sd = load_file(str(ckpt)) target.load_state_dict(sd, strict=True) # strict: any mismatch raises after = float(dict(target.named_parameters())[probe].detach().float().sum().item()) if before == after: raise RuntimeError(f"decoder weights UNCHANGED after load_state_dict (probe {probe} " f"sum {before}); refusing to run a checkpoint that did not apply.") target.eval() log(json.dumps({"full_ft_loaded": str(ckpt), "probe": probe, "sum_before": before, "sum_after": after, "tensors": len(sd)})) if lora_path: lora_status = dit_handler.add_lora(lora_path, adapter_name=adapter_name) log(lora_status) if not str(lora_status).startswith("✅"): raise RuntimeError(f"ACE LoRA load failed: {lora_status}") log(dit_handler.set_lora_scale(adapter_name, lora_scale)) log(dit_handler.set_use_lora(True)) llm_handler = None if use_lm or use_cot or not no_thinking: llm_handler = LLMHandler() lm_status, lm_success = llm_handler.initialize( checkpoint_dir=checkpoints, lm_model_path=lm_model, backend=lm_backend, device=device, offload_to_cpu=False, ) if not lm_success: raise RuntimeError(f"ACE LM init failed: {lm_status}") log(lm_status) return AceSession( dit_handler=dit_handler, llm_handler=llm_handler, device=device, lora_path=lora_path, lora_scale=lora_scale, full_ft_checkpoint=full_ft_checkpoint, load_seconds=round(time.time() - t0, 2), ) def build_params( meta: dict, item_dir: Path, *, task: str = "lego", steps: int = 64, seed: int = 1234, guidance_scale: float = 7.0, cover_strength: float = 0.45, no_thinking: bool = True, use_cot: bool = False, retake_seed: "int | None" = None, retake_variance: float = 0.0, ): """metadata.json -> (GenerationParams, GenerationConfig). ⚠️ THE SINGLE CONSTRUCTION. Every caller — CLI, resident worker, studio — comes through here. Do not copy this into another file; import it. See the module docstring for what a second copy cost. Values default to docs/INFERENCE_RECIPE.md (the May 7 recipe); metadata.json overrides where it carries a value. Note what the meta deliberately may omit: bpm absent means lego locks tempo from src_audio itself, which is correct — a GUESSED bpm is worse than none (the app sent a hardcoded 98 against audio at 119/170 and the model dutifully played out of tempo). """ from acestep.inference import GenerationParams, GenerationConfig role = meta.get("role", "") ace_role = "guitar" if role == "melody" else role # Recipe defaults, overridable by metadata. `or` not `.get(k, default)`: an empty # string in the metadata means "absent", not "deliberately empty" — and empty is # exactly the failure this guards. caption = (meta.get("ace_caption") or meta.get("prompt") or RECIPE_CAPTIONS.get(role) or f"Add {role} for this song.") global_caption = meta.get("global_caption") or RECIPE_GLOBAL_CAPTION source = item_dir / meta.get("source_audio", "context_mix_minus_target.wav") if task == "lego": instruction = f"Generate the {ace_role.upper()} track based on the audio context:" elif task == "complete": classes = meta.get("complete_track_classes") or [ace_role] instruction = "Complete the input track with " + " | ".join(str(c).upper() for c in classes) + ":" else: instruction = "Generate audio semantic tokens based on the given conditions:" params = GenerationParams( task_type=task, src_audio=str(source), instruction=instruction, caption=caption, global_caption=global_caption, lyrics=meta.get("lyrics", "[Instrumental]"), instrumental=bool(meta.get("instrumental", True)), vocal_language=meta.get("vocal_language", "unknown"), bpm=_meta_int(meta, "bpm"), keyscale=meta.get("keyscale") or meta.get("key") or "", timesignature=str(meta.get("timesignature") or meta.get("time_signature") or RECIPE_TIMESIGNATURE), duration=_meta_float(meta, "duration_seconds") or -1.0, repainting_start=0.0, repainting_end=-1, # Retake: variance-preserving variation. With a FIXED base `seed` per part and a # small `retake_variance`, each new sample is the SAME part subtly evolved — not # an unrelated diffusion draw. This is how a part 'continues from what it was'. retake_seed=retake_seed, retake_variance=float(retake_variance or 0.0), inference_steps=steps, seed=seed, thinking=not no_thinking, guidance_scale=guidance_scale, audio_cover_strength=cover_strength, use_cot_metas=use_cot, use_cot_caption=False, # production engine.py uses metas + language, NOT caption use_cot_language=use_cot, shift=3.0, # match production engine.py (was unset -> default) dcw_enabled=False, # #1255 NOISE FIX (production default) — the un-denoised culprit cover_noise_strength=0.0, use_adg=False, use_constrained_decoding=True, enable_normalization=True, normalization_db=-1.0, cfg_interval_start=0.0, cfg_interval_end=1.0, ) config = GenerationConfig(batch_size=1, use_random_seed=False, seeds=[seed], audio_format="wav") return params, config def run_take( session: AceSession, item_dir: Path, *, task: str = "lego", steps: int = 64, seed: int = 1234, guidance_scale: float = 7.0, cover_strength: float = 0.45, no_thinking: bool = True, use_cot: bool = False, out_dir: Path | None = None, generated_name: str | None = None, retake_seed: "int | None" = None, retake_variance: float = 0.0, log=print, ) -> dict: """One generation against an already-loaded session. Warm: ~7s. Identical to what the CLI does per item — because the CLI calls this. """ from acestep.inference import generate_music item_dir = Path(item_dir).resolve() meta = json.loads((item_dir / "metadata.json").read_text()) out_dir = Path(out_dir).resolve() if out_dir else item_dir / f"ace_{task}_corrected" out_dir.mkdir(parents=True, exist_ok=True) params, config = build_params( meta, item_dir, task=task, steps=steps, seed=seed, guidance_scale=guidance_scale, cover_strength=cover_strength, no_thinking=no_thinking, use_cot=use_cot, retake_seed=retake_seed, retake_variance=retake_variance, ) t0 = time.time() result = generate_music(session.dit_handler, session.llm_handler, params, config, save_dir=str(out_dir)) (out_dir / "result.json").write_text( json.dumps(result.to_dict() if hasattr(result, "to_dict") else result.__dict__, indent=2, default=str) + "\n" ) if not result.success: raise RuntimeError(result.error or result.status_message) generated = result.audios[0]["path"] target_name = generated_name or ("generated_full_ace.wav" if task == "cover" else "generated_stem_ace.wav") copied_to = item_dir / target_name shutil.copy2(generated, copied_to) return { "generated": generated, "copied_to": str(copied_to), "task": task, "lora_path": session.lora_path, "lora_scale": session.lora_scale, "seed": seed, "steps": steps, "guidance_scale": guidance_scale, "cover_strength": cover_strength, "generate_seconds": round(time.time() - t0, 2), } def main() -> None: p = argparse.ArgumentParser(description="Run ACE-Step corrected baseline task for one prepared pack item.") p.add_argument("item_dir", type=Path) p.add_argument("--task", choices=["lego", "complete", "cover"], default="lego") p.add_argument("--ace-root", default="/home/fcolo/ace-step-1.5-xl") p.add_argument("--checkpoints", default="/home/fcolo/ace-step/checkpoints") p.add_argument("--model", default="acestep-v15-xl-base") p.add_argument("--lm-model", default="acestep-5Hz-lm-1.7B") p.add_argument("--lm-backend", default="pt", choices=["pt", "vllm"]) p.add_argument("--device", default="cuda") p.add_argument("--steps", type=int, default=64) p.add_argument("--seed", type=int, default=1234) p.add_argument("--guidance-scale", type=float, default=7.0) p.add_argument("--cover-strength", type=float, default=0.45) p.add_argument("--no-thinking", action="store_true") p.add_argument("--use-lm", action="store_true", help="Initialize/use ACE 5Hz LM. Off by default for source-conditioned lego because CoT can rewrite the stem prompt.") p.add_argument("--use-cot", action="store_true", help="Allow LM CoT to rewrite/fill caption/language/metas. Off by default to preserve explicit stem prompts.") p.add_argument("--out-dir", type=Path, default=None, help="ACE raw output directory. Defaults to item_dir/ace__corrected.") p.add_argument("--generated-name", default=None, help="Name to copy the generated wav to in item_dir. Defaults to generated_stem_ace.wav or generated_full_ace.wav.") p.add_argument("--full-ft-checkpoint", default=None, help="Directory holding decoder_model.safetensors from a FULL fine-tune " "(e.g. artifacts/train/full_finetune_v4_.../best). Replaces the whole DiT " "decoder. Mutually exclusive with --lora-path: an adapter trained against " "the ORIGINAL decoder stacked on replaced weights is silently wrong, so " "passing both is refused rather than combined.") p.add_argument("--lora-path", default=None, help="Optional PEFT LoRA adapter directory to load after base ACE init.") p.add_argument("--adapter-name", default="stemgen", help="Adapter name used when loading --lora-path.") p.add_argument("--lora-scale", type=float, default=1.0) args = p.parse_args() if args.full_ft_checkpoint and args.lora_path: raise SystemExit("--full-ft-checkpoint and --lora-path are mutually exclusive") session = init_ace( ace_root=args.ace_root, checkpoints=args.checkpoints, model=args.model, device=args.device, lm_model=args.lm_model, lm_backend=args.lm_backend, no_thinking=args.no_thinking, use_lm=args.use_lm, use_cot=args.use_cot, full_ft_checkpoint=args.full_ft_checkpoint, lora_path=args.lora_path, adapter_name=args.adapter_name, lora_scale=args.lora_scale, ) out = run_take( session, args.item_dir, task=args.task, steps=args.steps, seed=args.seed, guidance_scale=args.guidance_scale, cover_strength=args.cover_strength, no_thinking=args.no_thinking, use_cot=args.use_cot, out_dir=args.out_dir, generated_name=args.generated_name, ) # Same shape the CLI has always printed. Callers parse this. print(json.dumps({k: out[k] for k in ("generated", "copied_to", "task", "lora_path", "lora_scale")}, indent=2)) if __name__ == "__main__": main()