File size: 1,277 Bytes
755da9f | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 | """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
|