Release unused host diffusion weights in preparation-only service

This commit is contained in:
firestar5683
2026-09-19 14:28:52 -07:00
parent ffa87f71ca
commit 6995c688db
3 changed files with 45 additions and 4 deletions
@@ -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
@@ -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()
+1 -1
View File
@@ -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}))