Download runtime/state_commit_audit.py from Accio-Lab/occamy-1.0-MTP: direct link, hf CLI and curl.
- Browser
- Download file 1.28 kB
-
https://huggingface.co/Accio-Lab/occamy-1.0-MTP/resolve/main/runtime/state_commit_audit.py
- Command line
-
hf download hf://Accio-Lab/occamy-1.0-MTP/runtime/state_commit_audit.py
-
curl -L -o state_commit_audit.py https://huggingface.co/Accio-Lab/occamy-1.0-MTP/resolve/main/runtime/state_commit_audit.py
1.28 kB
| """Validate the native single-request MTP state commit without changing it.""" | |
| import json | |
| from pathlib import Path | |
| import torch | |
| def install(config): | |
| from sglang.srt.speculative.eagle_worker_v2 import EAGLEWorkerV2 | |
| if getattr(EAGLEWorkerV2,'_strict_commit_audit',False):return | |
| EAGLEWorkerV2._strict_commit_audit=True | |
| original=EAGLEWorkerV2._mamba_verify_update | |
| path=Path(config['output']+'.commit') | |
| def checked(self,batch,lens,indices,bs): | |
| result=original(self,batch,lens,indices,bs) | |
| if bs==1 and not batch.forward_mode.is_idle(): | |
| backend=self.target_worker.model_runner.attn_backend.linear_attn_backend | |
| pool=backend.req_to_token_pool.get_speculative_mamba2_params_all_layers() | |
| slot=int(backend.forward_metadata.mamba_cache_indices[0]);step=int(lens[0])-1 | |
| ok=torch.equal(pool.temporal[:,slot],pool.intermediate_ssm[:,0,step]) and all(torch.equal(live[:,slot],buf[:,0,step]) for live,buf in zip(pool.conv,pool.intermediate_conv_window)) | |
| if not ok:raise RuntimeError('Strict MTP state commit mismatch') | |
| with path.open('a') as f:f.write(json.dumps({'accepted_nodes':step+1,'state_equal':ok})+'\n') | |
| return result | |
| EAGLEWorkerV2._mamba_verify_update=checked | |