From 6995c688db48e8651dec23a8e80be29565e8f24f Mon Sep 17 00:00:00 2001 From: firestar5683 <168790843+firestar5683@users.noreply.github.com> Date: Sat, 19 Sep 2026 14:28:52 -0700 Subject: [PATCH] Release unused host diffusion weights in preparation-only service --- .../host_hook_adapter.py | 28 +++++++++++++++++-- .../test_host_hook_adapter.py | 19 ++++++++++++- roadscore/prototype/hook_service.py | 2 +- 3 files changed, 45 insertions(+), 4 deletions(-) diff --git a/roadscore/experiments/ace_chestnut_20260916/host_hook_adapter.py b/roadscore/experiments/ace_chestnut_20260916/host_hook_adapter.py index 8c0f55515e..9ac86cd0a0 100644 --- a/roadscore/experiments/ace_chestnut_20260916/host_hook_adapter.py +++ b/roadscore/experiments/ace_chestnut_20260916/host_hook_adapter.py @@ -42,9 +42,24 @@ def continuation_inputs(context, prefix): return result, source, mask +class PreparationOnlyDecoder: + """Dispatch marker, never a model: accidental diffusion must fail closed.""" + def __call__(self, *args, **kwargs): + raise RuntimeError('Host diffusion disabled in preparation-only service') + + +def release_preparation_decoder(handler): + # Official code selects the capture boundary only when mlx_decoder is not None. + # Require real initialization first; do not fake a successful model conversion. + if not handler.use_mlx_dit or handler.mlx_decoder is None or handler.model.decoder is not None: + raise RuntimeError('Release requires verified MLX conversion and disabled Torch decoder') + handler.mlx_decoder = PreparationOnlyDecoder() + + class HostHookAdapter: """One serialized host model instance; initialize only on actual cache miss.""" - def __init__(self, assets_root): + def __init__(self, assets_root, *, preparation_only=False): + self.preparation_only = preparation_only self.base = Path(assets_root).resolve() / 'experiments/composition_20260916' self.handler = self.lm = None self.lock = threading.Lock() @@ -66,7 +81,7 @@ class HostHookAdapter: model = tree_hash(self.base / 'models/ace/checkpoints') official = tree_hash(self.base / 'vendor/ACE-Step-1.5/acestep') local = hashlib.sha256(Path(__file__).read_bytes()).hexdigest() - preparation = hashlib.sha256((official + local).encode()).hexdigest() + preparation = hashlib.sha256((official + local + ('preparation-only-v1' if self.preparation_only else 'full-host-v1')).encode()).hexdigest() self._fingerprints = {'model_fingerprint': model, 'preparation_fingerprint': preparation} return dict(self._fingerprints) @@ -107,6 +122,14 @@ class HostHookAdapter: device='mps', use_mlx_dit=True, offload_to_cpu=True, offload_dit_to_cpu=True) if not ok or not handler.use_mlx_dit or handler.mlx_decoder is None: raise RuntimeError(f'No safe MLX preparation boundary: {status}') + if self.preparation_only: + release_preparation_decoder(handler) + import gc + import mlx.core as mx + import torch + gc.collect() + mx.clear_cache() + torch.mps.empty_cache() self.handler, self.lm = handler, lm def __call__(self, request, sources, output): @@ -162,6 +185,7 @@ class HostHookAdapter: (output / 'prepared.json').write_text(json.dumps({ 'request_key': request.cache_key, 'semantic_seed': request.semantic_seed, 'semantic_plan_present': True, 'audio_diffusion_called': False, + 'host_preparation_only': self.preparation_only, 'conditioning_strategy': 'full semantic planning per bounded window; actual hook audio reference; committed latent prefix', 'native_validated': False, 'request': request.identity()}, indent=2)) from hook_planning import validate_prepared diff --git a/roadscore/experiments/ace_chestnut_20260916/test_host_hook_adapter.py b/roadscore/experiments/ace_chestnut_20260916/test_host_hook_adapter.py index f3f18d5aa2..4056b43ced 100644 --- a/roadscore/experiments/ace_chestnut_20260916/test_host_hook_adapter.py +++ b/roadscore/experiments/ace_chestnut_20260916/test_host_hook_adapter.py @@ -1,9 +1,10 @@ import unittest +import weakref import tempfile from pathlib import Path from types import SimpleNamespace import numpy as np -from host_hook_adapter import continuation_inputs, finish_capture +from host_hook_adapter import continuation_inputs, finish_capture, release_preparation_decoder, PreparationOnlyDecoder, HostHookAdapter from test_hook_planning import fake_prepare, request @@ -24,6 +25,22 @@ class PrefixTests(unittest.TestCase): with self.assertRaises(FileNotFoundError): finish_capture(True, expected, req, {}, output) + def test_preparation_only_releases_weights_but_keeps_capture_dispatch(self): + class Decoder: pass + handler = SimpleNamespace(use_mlx_dit=True, mlx_decoder=Decoder(), model=SimpleNamespace(decoder=None)) + reference = weakref.ref(handler.mlx_decoder) + release_preparation_decoder(handler) + self.assertIsNone(reference()) + self.assertIsInstance(handler.mlx_decoder, PreparationOnlyDecoder) + self.assertTrue(handler.use_mlx_dit) + with self.assertRaises(RuntimeError): handler.mlx_decoder(None) + + def test_release_cannot_mask_failed_conversion_or_torch_fallback(self): + for use_mlx, decoder, torch_decoder in [(False, object(), None), (True, None, None), (True, object(), object())]: + handler = SimpleNamespace(use_mlx_dit=use_mlx, mlx_decoder=decoder, model=SimpleNamespace(decoder=torch_decoder)) + with self.assertRaises(RuntimeError): release_preparation_decoder(handler) + self.assertFalse(HostHookAdapter('/unused').preparation_only) + def test_exactprefix_only_future_hints_survive(self): context = np.arange(1 * 1125 * 128, dtype=np.float32).reshape(1, 1125, 128) original = context.copy() diff --git a/roadscore/prototype/hook_service.py b/roadscore/prototype/hook_service.py index 48e0bc36d5..32f5ad502b 100644 --- a/roadscore/prototype/hook_service.py +++ b/roadscore/prototype/hook_service.py @@ -72,7 +72,7 @@ def main(): parser=argparse.ArgumentParser();parser.add_argument('--assets-root',type=Path,required=True);parser.add_argument('--cache',type=Path,required=True);parser.add_argument('--ready',type=Path,required=True);args=parser.parse_args() sys.path.insert(0,str(Path(__file__).resolve().parents[1]/'experiments/ace_chestnut_20260916')) from host_hook_adapter import HostHookAdapter - token=os.environ['ROADSCORE_PLANNER_TOKEN'];adapter=HostHookAdapter(args.assets_root) + token=os.environ['ROADSCORE_PLANNER_TOKEN'];adapter=HostHookAdapter(args.assets_root, preparation_only=True) server=make_server(adapter,PlanCache(args.cache),token) identity=adapter.fingerprints() args.ready.write_text(json.dumps({'port':server.server_port,**identity}))