File size: 19,123 Bytes
2f9dfba
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
#!/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_<task>_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()