Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
7 changes: 5 additions & 2 deletions acestep/engine/sa3_context.py
Original file line number Diff line number Diff line change
Expand Up @@ -187,9 +187,11 @@ def make_schedule_builder(
def _builder(denoise: float) -> torch.Tensor:
import stable_audio_3.inference.sampling as sampling

from .sa3_denoise_mapping import map_denoise_to_entry_sigma

schedule = sampling.build_schedule(
steps=int(steps),
sigma_max=float(denoise),
sigma_max=map_denoise_to_entry_sigma(float(denoise)),
dist_shift=prepared["dist_shift"],
effective_seq_len=prepared["effective_seq_len"],
fallback_seq_len=prepared["fallback_seq_len"],
Expand Down Expand Up @@ -254,7 +256,8 @@ class SA3SAMEWindowCodec:
Two execution paths, identical interface:

* **TRT** when ``use_trt`` and the built window engine exists
(``same_l_decode_window_t*``): ~9-10 ms per ~1 s window, latent
(``same_l_decode_window_<plugin_tag>_t*``): ~9-10 ms per ~1 s
window, latent
scaled by ``pretransform.scale`` before the call (spike
``scale_mode="pretransform"``, rel_rms ~8e-3 vs eager full).
* **Eager** fallback: the spike's ``decode_sa3_latent_window``
Expand Down
77 changes: 77 additions & 0 deletions acestep/engine/sa3_denoise_mapping.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,77 @@
"""Optional product mapping for the Stable Audio 3 denoise control.

The corrected upstream schedule makes ``sigma_max`` mathematically monotonic,
but measured audio change is still concentrated near the top of its range. In
a 41-point sweep over nine clips, less than one fifth of the available change
occurred below sigma 0.70. The knot table below inverts that measured curve so
equal control movements target roughly equal movements in the referee score.

The referee combines harmonic and rhythmic change and reproduced two separate
listening-order judgments that a loudness-based measure did not. It validates
the ordering and the location of the dead region, not a claim that every step
is equally perceptible. This is also one global mapping: sparse acoustic input
can move faster than dense electronic material.

Set ``DEMON_SA3_DENOISE_MAPPING=identity`` to bypass the product mapping while
retaining the upstream monotonic-schedule bugfix.
"""

from __future__ import annotations

import os
import typing as tp

__all__ = [
"denoise_mapping_mode",
"dial_to_entry_sigma",
"map_denoise_to_entry_sigma",
]


# Dial position -> entry sigma, obtained by inverting the measured change curve.
# Both endpoints remain exact: zero preserves the source and one starts from
# pure noise. Values between measured knots interpolate linearly.
_CALIBRATION: tp.Sequence[tuple[float, float]] = (
(0.00, 0.0000),
(0.10, 0.3786),
(0.20, 0.5762),
(0.30, 0.6853),
(0.40, 0.7321),
(0.50, 0.7683),
(0.65, 0.8376),
(0.80, 0.8843),
(1.00, 1.0000),
)


def denoise_mapping_mode() -> str:
"""Return ``calibrated`` (default) or the rollback mode ``identity``."""
mode = os.environ.get("DEMON_SA3_DENOISE_MAPPING", "calibrated").strip().lower()
if mode not in ("calibrated", "identity"):
raise ValueError(
"DEMON_SA3_DENOISE_MAPPING must be calibrated|identity, "
f"got {mode!r}"
)
return mode


def dial_to_entry_sigma(dial: float) -> float:
"""Interpolate the measured monotonic mapping, clamped to ``[0, 1]``."""
position = min(max(float(dial), 0.0), 1.0)
if position <= _CALIBRATION[0][0]:
return _CALIBRATION[0][1]

for (p0, s0), (p1, s1) in zip(_CALIBRATION, _CALIBRATION[1:]):
if position <= p1:
fraction = (position - p0) / (p1 - p0)
return s0 + (s1 - s0) * fraction

return _CALIBRATION[-1][1]


def map_denoise_to_entry_sigma(dial: float) -> float:
"""Apply the selected product mapping without changing endpoint semantics."""
clamped = min(max(float(dial), 0.0), 1.0)
if denoise_mapping_mode() == "identity":
return clamped
return dial_to_entry_sigma(clamped)
9 changes: 7 additions & 2 deletions acestep/engine/sa3_helpers.py
Original file line number Diff line number Diff line change
Expand Up @@ -25,7 +25,10 @@
# DEMON tracks this fork branch until the SA3 TensorRT/FP8 producer work
# merges upstream. The hash is the reproducibility boundary for installs.
SA3_VENDOR_URL = "https://github.com/ryanontheinside/stable-audio-3"
SA3_VENDOR_SHA = "03992120a3e296562271f209f4dd61dcfb55afff"
SA3_VENDOR_SHA = "960da1f8cbe205ab3b702edbfabd91113ab22473"
# Revision that last changed the source compiled into the SAME-L TensorRT
# engine. Keep this stable across vendor bumps that only touch other code.
SA3_SAME_L_PLUGIN_REVISION = "c07698548567fe6f163806f692d282bbaa57aba3"
SA3_VENDOR_ENV = "DEMON_SA3_SRC"
SA3_VENDOR_DIRNAME = "stable-audio-3"

Expand Down Expand Up @@ -152,7 +155,9 @@ def ensure_sa3_vendor(
reported as an error instead of being overwritten.
"""
vendor = sa3_vendor_dir()
if not (vendor / ".git").is_dir():
# A normal clone has a .git directory; a git worktree has a .git file
# pointing at the parent repository. Both are valid developer overrides.
if not (vendor / ".git").exists():
if check_only:
raise FileNotFoundError(f"SA3 vendor source is missing at {vendor}")
vendor.parent.mkdir(parents=True, exist_ok=True)
Expand Down
28 changes: 25 additions & 3 deletions acestep/engine/sa3_trt.py
Original file line number Diff line number Diff line change
Expand Up @@ -48,6 +48,7 @@
from __future__ import annotations

import math
import os
import re
import sys
import threading
Expand All @@ -58,7 +59,7 @@

from acestep import paths
from acestep.engine.obs import logger
from acestep.engine.sa3_helpers import sa3_vendor_dir
from acestep.engine.sa3_helpers import SA3_SAME_L_PLUGIN_REVISION, sa3_vendor_dir

IO_CHANNELS = 256
T5_TOKENS = 256
Expand Down Expand Up @@ -95,7 +96,10 @@
_DIT_REFIT_DIR_RE = re.compile(
r"^(?P<prefix>.+_dit)_refit_l(?P<lo>\d+)_(?P<opt>\d+)_(?P<hi>\d+)$"
)
_SAME_L_DIR_RE = re.compile(r"^same_l_decode_window_t(?P<lo>\d+)_(?P<opt>\d+)_(?P<hi>\d+)$")
_SAME_L_DIR_RE = re.compile(
r"^same_l_decode_window_(?P<tag>[a-z0-9_]+)_t"
r"(?P<lo>\d+)_(?P<opt>\d+)_(?P<hi>\d+)$"
)

# Deserialized-engine process cache. Engines are immutable post-load and
# support multiple execution contexts, so sharing one deserialization
Expand All @@ -108,6 +112,23 @@
_SAME_PLUGIN_REGISTERED = False


def same_l_plugin_build_tag() -> str:
"""Identity of the plugin implementation compiled into a SAME-L engine.

The upstream plugin is part of the serialized TensorRT engine, so changing
the vendored source revision or its selected AOT backend requires a new
engine even when the ONNX graph is unchanged.
"""
requested_plugin = os.environ.get("SA3_SWA_PLUGIN", "aot").strip().lower()
plugin = "jit" if requested_plugin == "jit" else "aot"
if plugin == "jit":
implementation = "jit"
else:
requested_backend = os.environ.get("SA3_SWA_AOT", "mma").strip().lower()
implementation = "mma" if requested_backend == "mma" else "ptx"
return f"{plugin}_{implementation}_v{SA3_SAME_L_PLUGIN_REVISION[:12]}"


def trt_engines_dir() -> Path:
return paths.models_dir() / "sa3" / "trt_engines"

Expand Down Expand Up @@ -286,9 +307,10 @@ def find_same_l_window_engine() -> Optional[tuple]:
base = trt_engines_dir()
if not base.is_dir():
return None
expected_tag = same_l_plugin_build_tag()
for sub in base.iterdir():
m = _SAME_L_DIR_RE.match(sub.name)
if not m:
if not m or m.group("tag") != expected_tag:
continue
f = sub / f"{sub.name}.trt"
if f.is_file():
Expand Down
21 changes: 17 additions & 4 deletions acestep/engine/trt/sa3_build.py
Original file line number Diff line number Diff line change
Expand Up @@ -87,6 +87,7 @@
SAMPLES_PER_LATENT,
T5_TOKENS,
_register_same_plugin,
same_l_plugin_build_tag,
trt_engines_dir,
)

Expand Down Expand Up @@ -218,10 +219,11 @@ class SameLWindowBuildConfig:
max_latents: int
workspace_gb: float = 16.0
onnx_files: list[str] = field(default_factory=lambda: list(SAME_L_ONNX_FILES))
plugin_build_tag: str = field(default_factory=same_l_plugin_build_tag)

def engine_name(self) -> str:
return (
f"same_l_decode_window_t{self.min_latents}"
f"same_l_decode_window_{self.plugin_build_tag}_t{self.min_latents}"
f"_{self.opt_latents}_{self.max_latents}"
)

Expand Down Expand Up @@ -292,6 +294,7 @@ def _build_strongly_typed_engine(
workspace_gb: float,
profile_shapes: dict[str, tuple[tuple, tuple, tuple]],
refit: bool = False,
python_plugin_preference: str | None = None,
) -> None:
"""Parse + build one STRONGLY_TYPED engine and serialize it to disk.

Expand All @@ -304,9 +307,16 @@ def _build_strongly_typed_engine(

trt_logger = trt.Logger(trt.Logger.WARNING)
builder = trt.Builder(trt_logger)
network = builder.create_network(
1 << int(trt.NetworkDefinitionCreationFlag.STRONGLY_TYPED)
)
network_flags = 1 << int(trt.NetworkDefinitionCreationFlag.STRONGLY_TYPED)
if python_plugin_preference == "aot":
network_flags |= 1 << int(
trt.NetworkDefinitionCreationFlag.PREFER_AOT_PYTHON_PLUGINS
)
elif python_plugin_preference == "jit":
network_flags |= 1 << int(
trt.NetworkDefinitionCreationFlag.PREFER_JIT_PYTHON_PLUGINS
)
network = builder.create_network(network_flags)
parser = trt.OnnxParser(network, trt_logger)
if not parser.parse_from_file(onnx_path):
for i in range(parser.num_errors):
Expand Down Expand Up @@ -494,6 +504,9 @@ def _build_same_l_window_engine(
profile_shapes={
"latent": ((1, IO_CHANNELS, lo), (1, IO_CHANNELS, opt), (1, IO_CHANNELS, hi)),
},
python_plugin_preference=(
"jit" if config.plugin_build_tag.startswith("jit_") else "aot"
),
)
_write_metadata(engine_path=engine_path, expected=expected, env=env)
elapsed = time.time() - t0
Expand Down
9 changes: 5 additions & 4 deletions acestep/streaming/sa3_backend.py
Original file line number Diff line number Diff line change
Expand Up @@ -143,10 +143,11 @@ def sa3_knob_specs(loras: tuple | list = ()) -> list:
KnobSpec(
"sa3_denoise", default=1.0, max_val=1.0, group="sa3",
description=(
"SA3 init_noise_level: fresh-noise vs source-anchor mix "
"at slot init (1.0 = generate from pure noise, lower = "
"closer cover of the source). Distinct from ACE's "
"'denoise' (k1 strength), hence the prefix."
"Measured SA3 audio-change amount: 1.0 generates from pure "
"noise, while lower values stay progressively closer to the "
"source. Mapped onto init_noise_level so useful change is "
"spread across the dial. Distinct from ACE's 'denoise' "
"(k1 strength), hence the prefix."
),
),
KnobSpec(
Expand Down
53 changes: 53 additions & 0 deletions tests/unit/test_sa3_denoise_mapping.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,53 @@
import pytest

from acestep.engine.sa3_denoise_mapping import (
denoise_mapping_mode,
dial_to_entry_sigma,
map_denoise_to_entry_sigma,
)


def test_mapping_is_monotonic_and_preserves_endpoints():
values = [dial_to_entry_sigma(i / 100) for i in range(101)]

assert values[0] == 0.0
assert values[-1] == 1.0
assert all(a <= b for a, b in zip(values, values[1:]))


def test_mapping_matches_measured_knots():
assert dial_to_entry_sigma(0.1) == pytest.approx(0.3786)
assert dial_to_entry_sigma(0.5) == pytest.approx(0.7683)
assert dial_to_entry_sigma(0.8) == pytest.approx(0.8843)


def test_mapping_interpolates_between_knots():
midpoint = (0.7683 + 0.8376) / 2
assert dial_to_entry_sigma(0.575) == pytest.approx(midpoint)


def test_mapping_clamps_inputs():
assert dial_to_entry_sigma(-1) == 0.0
assert dial_to_entry_sigma(2) == 1.0


def test_default_mapping_lifts_the_measured_dead_region(monkeypatch):
monkeypatch.delenv("DEMON_SA3_DENOISE_MAPPING", raising=False)

assert denoise_mapping_mode() == "calibrated"
assert map_denoise_to_entry_sigma(0.2) > 0.5
assert map_denoise_to_entry_sigma(0.5) > 0.7


def test_identity_mode_is_a_bugfix_preserving_rollback(monkeypatch):
monkeypatch.setenv("DEMON_SA3_DENOISE_MAPPING", "identity")

assert map_denoise_to_entry_sigma(0.2) == 0.2
assert map_denoise_to_entry_sigma(0.5) == 0.5


def test_invalid_mode_fails_loudly(monkeypatch):
monkeypatch.setenv("DEMON_SA3_DENOISE_MAPPING", "mystery")

with pytest.raises(ValueError, match=r"calibrated\|identity"):
map_denoise_to_entry_sigma(0.5)
60 changes: 60 additions & 0 deletions tests/unit/test_sa3_trt_engine_identity.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,60 @@
from pathlib import Path

from acestep.engine import sa3_trt
from acestep.engine.trt.sa3_build import SameLWindowBuildConfig


def _engine(root: Path, name: str) -> Path:
path = root / "sa3" / "trt_engines" / name / f"{name}.trt"
path.parent.mkdir(parents=True, exist_ok=True)
path.write_bytes(b"engine")
return path


def test_same_l_engine_name_carries_plugin_build_identity(monkeypatch):
monkeypatch.delenv("SA3_SWA_AOT", raising=False)
config = SameLWindowBuildConfig(32, 56, 96)

assert config.plugin_build_tag == sa3_trt.same_l_plugin_build_tag()
assert config.engine_name() == (
f"same_l_decode_window_{config.plugin_build_tag}_t32_56_96"
)


def test_same_l_discovery_ignores_legacy_engine(monkeypatch, tmp_path):
monkeypatch.setenv("ACESTEP_MODELS_DIR", str(tmp_path))
monkeypatch.delenv("SA3_SWA_AOT", raising=False)
_engine(tmp_path, "same_l_decode_window_t32_56_96")

assert sa3_trt.find_same_l_window_engine() is None


def test_same_l_discovery_selects_current_plugin_engine(monkeypatch, tmp_path):
monkeypatch.setenv("ACESTEP_MODELS_DIR", str(tmp_path))
monkeypatch.delenv("SA3_SWA_AOT", raising=False)
tag = sa3_trt.same_l_plugin_build_tag()
expected = _engine(tmp_path, f"same_l_decode_window_{tag}_t32_56_96")

assert sa3_trt.find_same_l_window_engine() == (expected, 32, 96)


def test_same_l_aot_backend_is_part_of_identity(monkeypatch):
monkeypatch.setenv("SA3_SWA_PLUGIN", "aot")
monkeypatch.setenv("SA3_SWA_AOT", "mma")
mma = sa3_trt.same_l_plugin_build_tag()
monkeypatch.setenv("SA3_SWA_AOT", "ptx")
ptx = sa3_trt.same_l_plugin_build_tag()

assert mma != ptx
assert mma.startswith("aot_mma_v")
assert ptx.startswith("aot_ptx_v")


def test_same_l_plugin_kind_is_part_of_identity(monkeypatch):
monkeypatch.setenv("SA3_SWA_PLUGIN", "aot")
aot = sa3_trt.same_l_plugin_build_tag()
monkeypatch.setenv("SA3_SWA_PLUGIN", "jit")
jit = sa3_trt.same_l_plugin_build_tag()

assert aot != jit
assert jit.startswith("jit_jit_v")
Loading