"""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