Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
18 commits
Select commit Hold shift + click to select a range
4cc120b
feat(inference): add ComponentDeviceMap for multi-GPU placement (PR1)
Steve Jul 6, 2026
a38527f
feat(inference): add multi-GPU auto-layout and cross-GPU routing (PR2)
Steve Jul 6, 2026
83c3c0e
refactor(device_map): split package to satisfy 200-LOC module cap.
Steve Jul 6, 2026
7d2b302
fix(inference): complete cross-GPU diffusion routing and APG device i…
Steve Jul 6, 2026
18c8d4a
fix(device_map): reserve DiT VRAM for co-located LM auto-layout fallb…
Steve Jul 6, 2026
c1d9d98
fix: address CodeRabbit parsing dead code and APG MPS regression.
Steve Jul 6, 2026
b6d4061
Fix silence_latent device check for multi-GPU CUDA indices.
Steve Jul 6, 2026
0c21299
Hoist normalize_component_device import and narrow device errors.
Steve Jul 6, 2026
f7ba683
Catch specific exceptions in device string parsing helpers.
Steve Jul 6, 2026
8f1463f
Fix VAE VRAM preflight to use mapped component device index.
Steve Jul 6, 2026
1d19b9b
Raise DeviceMapError for malformed CUDA device indices.
Steve Jul 6, 2026
c3ef2df
Split generate_music_decode tests to satisfy 200 LOC cap.
Steve Jul 6, 2026
7aee60f
Use normal imports in generate_music_decode test support.
Steve Jul 6, 2026
ec02094
fix(inference): honor mapped cuda:N for LM init and Gradio wiring
Steve Jul 11, 2026
a42ff62
Add multi-GPU CLI/API flags, status fields, and docs (PR3).
Steve Jul 6, 2026
1eba94f
fix(device_map): address review feedback on layout and discovery.
Steve Jul 6, 2026
216d3b2
refactor: address CodeRabbit module-size and VRAM review notes
Steve Jul 12, 2026
7085b67
fix(cli): match LM size token as -4B, not substring 4B
Steve Jul 12, 2026
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
657 changes: 61 additions & 596 deletions acestep/acestep_v15_pipeline.py

Large diffs are not rendered by default.

150 changes: 150 additions & 0 deletions acestep/acestep_v15_pipeline_gpu_mapping_test.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,150 @@
"""Unit tests for multi-GPU Gradio CLI flags in the pipeline."""

from __future__ import annotations

import os
import sys
import unittest
from types import SimpleNamespace
from unittest.mock import MagicMock, patch

from acestep import acestep_v15_pipeline


class PipelineGpuMappingTests(unittest.TestCase):
"""Verify multi-GPU CLI flags are wired into service initialization."""

def test_list_gpus_exits_after_printing_inventory(self) -> None:
"""``--list-gpus`` should print inventory and exit without launching UI."""
with patch.object(sys, "argv", ["acestep", "--list-gpus"]), patch(
"acestep.gradio_pipeline_cli.format_gpu_list_text",
return_value="GPU TABLE",
) as mock_format, patch(
"acestep.gradio_pipeline_cli.sys.exit",
side_effect=SystemExit(0),
) as mock_exit, patch(
"acestep.acestep_v15_pipeline.get_gpu_config",
return_value=SimpleNamespace(
gpu_memory_gb=24.0,
tier="tier6b",
max_duration_with_lm=480,
max_duration_without_lm=600,
max_batch_size_with_lm=8,
max_batch_size_without_lm=8,
init_lm_default=True,
available_lm_models=["acestep-5Hz-lm-0.6B"],
recommended_backend="vllm",
lm_backend_restriction=None,
offload_dit_to_cpu_default=False,
quantization_default=False,
),
), patch(
"acestep.acestep_v15_pipeline.set_global_gpu_config"
), patch(
"acestep.acestep_v15_pipeline.is_mps_platform",
return_value=False,
), patch(
"acestep.acestep_v15_pipeline.get_i18n"
), patch(
"acestep.gradio_pipeline_cli.available_languages_info",
return_value=[("en", "English", "English")],
), patch(
"acestep.acestep_v15_pipeline.os.makedirs"
):
with self.assertRaises(SystemExit):
acestep_v15_pipeline.main()
mock_format.assert_called_once()
mock_exit.assert_called_once_with(0)

def test_gpu_mapping_passed_to_initialize_service(self) -> None:
"""``--gpu-mapping`` must reach DiT init and drive LM device selection."""
gpu_config = SimpleNamespace(
gpu_memory_gb=24.0,
tier="tier6b",
max_duration_with_lm=480,
max_duration_without_lm=600,
max_batch_size_with_lm=8,
max_batch_size_without_lm=8,
init_lm_default=True,
available_lm_models=["acestep-5Hz-lm-0.6B"],
recommended_backend="vllm",
lm_backend_restriction=None,
offload_dit_to_cpu_default=False,
quantization_default=False,
)
dit_handler = MagicMock()
dit_handler.get_available_acestep_v15_models.return_value = ["acestep-v15-turbo"]
dit_handler.is_flash_attention_available.return_value = False
dit_handler.initialize_service.return_value = ("ok", True)
dit_handler.device_map = SimpleNamespace(lm="cuda:1")

llm_handler = MagicMock()
llm_handler.get_available_5hz_lm_models.return_value = ["acestep-5Hz-lm-0.6B"]
llm_handler.initialize.return_value = ("ok", True)

demo = MagicMock()
demo.queue.return_value = demo
demo.launch.return_value = None
captured: dict[str, object] = {}

def _create_demo(init_params=None, language="en"):
"""Capture init_params while returning a stub Gradio demo."""
captured["init_params"] = init_params
return demo

with patch.object(
sys,
"argv",
[
"acestep",
"--init_service",
"true",
"--init_llm",
"true",
"--config_path",
"acestep-v15-turbo",
"--gpu-mapping",
"auto",
],
), patch.dict(os.environ, {}, clear=True), patch(
"acestep.acestep_v15_pipeline.get_gpu_config",
return_value=gpu_config,
), patch(
"acestep.acestep_v15_pipeline.set_global_gpu_config"
), patch(
"acestep.acestep_v15_pipeline.is_mps_platform",
return_value=False,
), patch(
"acestep.acestep_v15_pipeline.get_i18n"
), patch(
"acestep.gradio_pipeline_cli.available_languages_info",
return_value=[("en", "English", "English")],
), patch(
"acestep.gradio_pipeline_startup.AceStepHandler",
return_value=dit_handler,
), patch(
"acestep.gradio_pipeline_startup.LLMHandler",
return_value=llm_handler,
), patch(
"acestep.acestep_v15_pipeline.create_demo",
side_effect=_create_demo,
), patch(
"acestep.gradio_pipeline_startup.ensure_lm_model",
return_value=(True, "ok"),
), patch(
"acestep.acestep_v15_pipeline.os.makedirs"
), patch(
"acestep.gradio_pipeline_startup.log_lm_device_deprecation"
):
acestep_v15_pipeline.main()

self.assertEqual(
"auto",
dit_handler.initialize_service.call_args.kwargs["gpu_mapping"],
)
self.assertEqual("cuda:1", llm_handler.initialize.call_args.kwargs["device"])
self.assertEqual("auto", captured["init_params"]["gpu_mapping"])


if __name__ == "__main__":
unittest.main()
28 changes: 13 additions & 15 deletions acestep/acestep_v15_pipeline_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -37,7 +37,7 @@ def _run_main(
argv: list[str],
*,
env: dict[str, str] | None = None,
) -> tuple[MagicMock, dict[str, object]]:
) -> tuple[MagicMock, MagicMock, dict[str, object]]:
"""Run ``main`` with heavy dependencies stubbed and capture startup state."""
gpu_config = self._legacy_gpu_config()
dit_handler = MagicMock()
Expand All @@ -52,15 +52,17 @@ def _run_main(
demo = MagicMock()
demo.queue.return_value = demo
demo.launch.return_value = None

captured: dict[str, object] = {}

def _create_demo(init_params=None, language="en"):
"""Capture init_params while returning a stub Gradio demo."""
captured["init_params"] = init_params
captured["language"] = language
return demo

with patch.object(sys, "argv", argv), patch.dict(os.environ, env or {}, clear=True), patch(
with patch.object(sys, "argv", argv), patch.dict(
os.environ, env or {}, clear=True
), patch(
"acestep.acestep_v15_pipeline.get_gpu_config",
return_value=gpu_config,
), patch(
Expand All @@ -71,30 +73,30 @@ def _create_demo(init_params=None, language="en"):
), patch(
"acestep.acestep_v15_pipeline.get_i18n"
), patch(
"acestep.acestep_v15_pipeline.available_languages_info",
"acestep.gradio_pipeline_cli.available_languages_info",
return_value=[("en", "English", "English")],
), patch(
"acestep.acestep_v15_pipeline.AceStepHandler",
"acestep.gradio_pipeline_startup.AceStepHandler",
return_value=dit_handler,
), patch(
"acestep.acestep_v15_pipeline.LLMHandler",
"acestep.gradio_pipeline_startup.LLMHandler",
return_value=llm_handler,
), patch(
"acestep.acestep_v15_pipeline.create_demo",
side_effect=_create_demo,
), patch(
"acestep.acestep_v15_pipeline.ensure_lm_model",
"acestep.gradio_pipeline_startup.ensure_lm_model",
return_value=(True, "ok"),
), patch(
"acestep.acestep_v15_pipeline.os.makedirs"
):
acestep_v15_pipeline.main()

return llm_handler, captured
return llm_handler, dit_handler, captured

def test_main_forces_pt_backend_for_explicit_vllm_argument(self) -> None:
"""Legacy CUDA startup should override an explicit CLI vLLM request."""
llm_handler, captured = self._run_main(
llm_handler, _, captured = self._run_main(
[
"acestep",
"--init_service",
Expand All @@ -109,29 +111,26 @@ def test_main_forces_pt_backend_for_explicit_vllm_argument(self) -> None:
"vllm",
]
)

self.assertEqual("pt", llm_handler.initialize.call_args.kwargs["backend"])
self.assertEqual("pt", captured["init_params"]["backend"])

def test_main_forces_pt_backend_for_service_mode_backend_override(self) -> None:
"""Service mode should not re-enable vLLM on legacy CUDA hardware."""
llm_handler, captured = self._run_main(
llm_handler, _, captured = self._run_main(
["acestep", "--service_mode", "true", "--init_llm", "true"],
env={"SERVICE_MODE_BACKEND": "vllm"},
)

self.assertEqual("pt", llm_handler.initialize.call_args.kwargs["backend"])
self.assertEqual("pt", captured["init_params"]["backend"])

def test_main_forces_pt_backend_for_api_env_override(self) -> None:
"""API-mode env overrides should still resolve to the safe startup backend."""
api_routes_module = types.SimpleNamespace(setup_api_routes=MagicMock())

with patch.dict(
sys.modules,
{"acestep.ui.gradio.api.api_routes": api_routes_module},
), patch("time.sleep", side_effect=KeyboardInterrupt):
llm_handler, captured = self._run_main(
llm_handler, _, captured = self._run_main(
[
"acestep",
"--enable-api",
Expand All @@ -144,7 +143,6 @@ def test_main_forces_pt_backend_for_api_env_override(self) -> None:
],
env={"ACESTEP_LM_BACKEND": "vllm"},
)

self.assertEqual("pt", llm_handler.initialize.call_args.kwargs["backend"])
self.assertEqual("pt", captured["init_params"]["backend"])

Expand Down
12 changes: 12 additions & 0 deletions acestep/api/http/model_init_service.py
Original file line number Diff line number Diff line change
Expand Up @@ -117,6 +117,7 @@ def initialize_models_for_request(
compile_model=compile_model,
offload_to_cpu=offload_to_cpu,
offload_dit_to_cpu=offload_dit_to_cpu,
gpu_mapping=os.getenv("ACESTEP_GPU_MAPPING"),
)
if not ok:
setattr(app_state, error_attr, status_msg)
Expand All @@ -137,6 +138,17 @@ def initialize_models_for_request(

lm_backend = resolve_lm_backend(os.getenv("ACESTEP_LM_BACKEND"), gpu_config)
lm_device = os.getenv("ACESTEP_LM_DEVICE", device)
device_map = getattr(handler, "device_map", None)
using_device_map_lm = False
if device_map is not None and device_map.lm is not None:
lm_device = device_map.lm
using_device_map_lm = True
from acestep.device_map import log_lm_device_deprecation

log_lm_device_deprecation(
explicit_lm_device=os.getenv("ACESTEP_LM_DEVICE"),
using_device_map_lm=using_device_map_lm,
)
lm_offload_env = os.getenv("ACESTEP_LM_OFFLOAD_TO_CPU")
lm_offload = env_bool("ACESTEP_LM_OFFLOAD_TO_CPU", False) if lm_offload_env is not None else offload_to_cpu

Expand Down
5 changes: 5 additions & 0 deletions acestep/api/http/model_service_routes.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,7 @@

from acestep.api.http.model_init_service import initialize_models_for_request
from acestep.constants import TASK_TYPES_BASE, TASK_TYPES_TURBO
from acestep.device_map import collect_gpu_runtime_status


class InitModelRequest(BaseModel):
Expand Down Expand Up @@ -122,6 +123,7 @@ def _collect_model_inventory(
"lm_models": lm_models,
"loaded_lm_model": loaded_lm_model,
"llm_initialized": llm_initialized,
**collect_gpu_runtime_status(getattr(app.state, "handler", None)),
}


Expand Down Expand Up @@ -157,6 +159,9 @@ async def health_check():
"llm_initialized": inventory["llm_initialized"],
"loaded_model": inventory["default_model"],
"loaded_lm_model": inventory["loaded_lm_model"],
"gpus": inventory["gpus"],
"gpu_mapping": inventory["gpu_mapping"],
"device_map": inventory["device_map"],
}
)

Expand Down
16 changes: 16 additions & 0 deletions acestep/api/http/model_service_routes_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -103,6 +103,22 @@ def test_collect_model_inventory_merges_loaded_and_available_models(self):
self.assertIn("acestep-v15-turbo", names)
self.assertEqual("acestep-v15-base", inventory["default_model"])
self.assertTrue(inventory["llm_initialized"])
self.assertIn("gpus", inventory)
self.assertIn("gpu_mapping", inventory)
self.assertIn("device_map", inventory)

def test_health_route_includes_gpu_runtime_fields(self):
"""Health endpoint should expose GPU inventory and mapping metadata."""

app = self._build_app()
endpoint = _get_endpoint(app, "/health", "GET")
with mock.patch("acestep.api.http.model_service_routes.os.path.isdir", return_value=False):
result = asyncio.run(endpoint())

self.assertEqual(200, result["code"])
self.assertIn("gpus", result["data"])
self.assertIn("gpu_mapping", result["data"])
self.assertIn("device_map", result["data"])

def test_init_route_wraps_initializer_exception(self):
"""Init endpoint should convert initializer exceptions into wrapped code=500 payloads."""
Expand Down
12 changes: 12 additions & 0 deletions acestep/api/startup_llm_init.py
Original file line number Diff line number Diff line change
Expand Up @@ -23,6 +23,7 @@ def initialize_llm_at_startup(
get_model_name: Callable[[str], str],
ensure_model_downloaded: Callable[[str, str], str],
env_bool: Callable[[str, bool], bool],
dit_handler: Any = None,
) -> None:
"""Initialize LLM model according to GPU config and environment overrides."""

Expand Down Expand Up @@ -73,6 +74,17 @@ def initialize_llm_at_startup(

lm_backend = resolve_lm_backend(os.getenv("ACESTEP_LM_BACKEND"), gpu_config)
lm_device = os.getenv("ACESTEP_LM_DEVICE", device)
device_map = getattr(dit_handler, "device_map", None) if dit_handler is not None else None
using_device_map_lm = False
if device_map is not None and device_map.lm is not None:
lm_device = device_map.lm
using_device_map_lm = True
from acestep.device_map import log_lm_device_deprecation

log_lm_device_deprecation(
explicit_lm_device=os.getenv("ACESTEP_LM_DEVICE"),
using_device_map_lm=using_device_map_lm,
)
lm_offload_env = os.getenv("ACESTEP_LM_OFFLOAD_TO_CPU")
lm_offload = env_bool("ACESTEP_LM_OFFLOAD_TO_CPU", False) if lm_offload_env is not None else offload_to_cpu

Expand Down
2 changes: 2 additions & 0 deletions acestep/api/startup_model_init.py
Original file line number Diff line number Diff line change
Expand Up @@ -85,6 +85,7 @@ def do_model_initialization(
compile_model=compile_model,
offload_to_cpu=offload_to_cpu,
offload_dit_to_cpu=offload_dit_to_cpu,
gpu_mapping=os.getenv("ACESTEP_GPU_MAPPING"),
)
if not ok:
app.state._init_error = status_msg
Expand Down Expand Up @@ -157,6 +158,7 @@ def do_model_initialization(
get_model_name=get_model_name,
ensure_model_downloaded=ensure_model_downloaded,
env_bool=env_bool,
dit_handler=handler,
)

print("[API Server] All models initialized successfully!")
Expand Down
9 changes: 7 additions & 2 deletions acestep/core/generation/handler/audio_codes.py
Original file line number Diff line number Diff line change
Expand Up @@ -56,7 +56,8 @@ def _decode_audio_codes_to_latents(self, code_str: str) -> Optional[torch.Tensor
with self._load_model_context("model"):
quantizer = self.model.tokenizer.quantizer
detokenizer = self.model.detokenizer
indices = torch.tensor(code_ids, device=self.device, dtype=torch.long)
dit_device = self._get_component_device("model")
indices = torch.tensor(code_ids, device=dit_device, dtype=torch.long)
indices = indices.unsqueeze(0).unsqueeze(-1)

quantized = quantizer.get_output_from_indices(indices)
Expand All @@ -83,7 +84,11 @@ def convert_src_audio_to_codes(self, audio_file) -> str:
return "❌ Audio file appears to be silent"
latents = self._encode_audio_to_latents(processed_audio)

attention_mask = torch.ones(latents.shape[0], dtype=torch.bool, device=self.device)
attention_mask = torch.ones(
latents.shape[0],
dtype=torch.bool,
device=self._get_component_device("model"),
)
with self._load_model_context("model"):
hidden_states = latents.unsqueeze(0)
_, indices, _ = self.model.tokenize(
Expand Down
Loading