sglang
mtp
speculative-decoding
draft-head
qwen3_5_moe
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