botp
/

File size: 9,209 Bytes
1d2de8a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
"""The confidence layer text reasoning on immutable answer-contract base-prefix states."""
import copy,hashlib,json
from pathlib import Path
from solomon.engine_numerics import CudaEngine

FORCE='\n</think>\n\nGive only the answer letter: '

def validate_request(state,n_letters,max_tokens):
    if not isinstance(max_tokens,int) or isinstance(max_tokens,bool) or not 1<=max_tokens<=3072:raise ValueError('reasoning cap must be 1..3072')
    if not isinstance(n_letters,int) or isinstance(n_letters,bool) or not 2<=n_letters<=10:raise ValueError('invalid letter count')
    if any('image' in p for p in state['parts']):raise ValueError('the confidence layer reasoning confidence supports text only')

class ReasoningEngine(CudaEngine):
    def reasoning_identity(self):
        return {'base_runtime':self.identity,'reasoning_code_sha256':hashlib.sha256(Path(__file__).read_bytes()).hexdigest(),'adapter':'disabled','forced_suffix':FORCE}

    def estimate(self,state,block,n,stage='fast',max_new_tokens=512):
        if stage not in ('fast','reasoning'):raise ValueError('unknown stage')
        if any('image' in p for p in state['parts']):raise ValueError('the confidence layer text-only budget estimator')
        text=self._render(state['parts'],block)
        if stage=='reasoning':
            validate_request(state,n,max_new_tokens)
            marker='<|im_start|>assistant\n'
            if marker not in text:raise ValueError('unknown assistant template')
            text=text[:text.rfind(marker)]+marker+'<think>\n'
        ids=self.t.encode(text,add_special_tokens=False)
        prefix=state['prefix_ids'][0].tolist()
        if ids[:len(prefix)]!=prefix:raise ValueError('budget prefix mismatch')
        extra=len(self.t.encode(FORCE,add_special_tokens=False)) if stage=='reasoning' else 0
        return {'input_tokens':len(ids)-len(prefix)+extra,'generated_tokens':max_new_tokens if stage=='reasoning' else 0,'branches':1,'full_prompt_tokens':len(ids),'reused_prefix_tokens':len(prefix)}

    def _reason_inputs(self,state,block):
        text=self._render(state['parts'],block)
        marker='<|im_start|>assistant\n'
        if marker not in text:raise ValueError('unknown assistant template')
        text=text[:text.rfind(marker)]+marker+'<think>\n'
        ids,embeds,pos,_=self._encode(state['parts'],text,state['features'])
        if not self.torch.equal(ids[:,:state['prefix_tokens']],state['prefix_ids']):raise ValueError('reasoning prefix mismatch')
        return text,ids,embeds,pos

    def reason(self,state,block,n_letters,max_tokens=512,execution='cached',replay_tokens=None,cap_checkpoints=None):
        """Greedy base reasoning then bounded forced letter readout; replay is qualification only."""
        validate_request(state,n_letters,max_tokens)
        checkpoints=tuple(cap_checkpoints or ())
        if any(not isinstance(c,int) or isinstance(c,bool) or not 1<=c<=max_tokens for c in checkpoints):raise ValueError('invalid cap checkpoints')
        if checkpoints and (replay_tokens is not None or execution!='cached'):raise ValueError('checkpoints require cached generation')
        if self.placement!='question':raise ValueError('reasoning needs base document state')
        if execution not in ('cached','full'):raise ValueError('bad execution')
        torch=self.torch;p=state['prefix_tokens'];previous=self.ctx['start']
        from transformers import DynamicCache
        try:
            self.ctx['start']=None
            with torch.inference_mode():
                text,ids,embeds,pos=self._reason_inputs(state,block)
                cache=copy.deepcopy(state['cache']) if execution=='cached' else DynamicCache(config=self.lm.config)
                start=p if execution=='cached' else 0
                h=self.lm(inputs_embeds=embeds[:,start:],position_ids=pos[...,start:],past_key_values=cache,use_cache=True).last_hidden_state
                generated=[];snapshots={};stop='token_cap';lastpos=pos[...,-1:];close=self.t.encode('</think>',add_special_tokens=False)
                eos=self.t.eos_token_id
                for i in range(0 if execution=='full' and replay_tokens is not None else max_tokens):
                    token=int(replay_tokens[i]) if replay_tokens is not None and i<len(replay_tokens) else int(self.model.lm_head(h[:,-1]).argmax(-1))
                    if replay_tokens is not None and i>=len(replay_tokens):break
                    generated.append(token)
                    one=torch.tensor([[token]],device=ids.device);lastpos=lastpos+1
                    h=self.lm(input_ids=one,position_ids=lastpos,past_key_values=cache,use_cache=True).last_hidden_state
                    if len(generated) in checkpoints:
                        sf=self.t.encode(FORCE,add_special_tokens=False);sfids=torch.tensor([sf],device=ids.device)
                        sfpos=lastpos+torch.arange(1,len(sf)+1,device=ids.device).view(1,1,-1)
                        sh=self.lm(input_ids=sfids,position_ids=sfpos,past_key_values=copy.deepcopy(cache),use_cache=True).last_hidden_state
                        sl=self.model.lm_head(sh[:,-1]).float()[0]
                        li=[self.t.encode(c,add_special_tokens=False)[0] for c in 'ABCDEFGHIJ'[:n_letters]]
                        sc=sl[li];sp=sc.softmax(-1)
                        snapshots[str(len(generated))]={'letter_logits':sc.cpu().tolist(),'probabilities':sp.cpu().tolist(),'prediction':int(sp.argmax()),'mass':float(sl.softmax(-1)[li].sum()),'generation_tokens':len(generated),'finish_reason':'token_cap'}
                    if replay_tokens is None and (token==eos or generated[-len(close):]==close):stop='closed_reasoning' if token!=eos else 'eos';break
                if execution=='full' and replay_tokens is not None:
                    generated=list(replay_tokens[:max_tokens]);lastpos=pos[...,-1:]+len(generated)
                suffix=self.t.encode(FORCE,add_special_tokens=False);suffix_ids=torch.tensor([suffix],device=ids.device)
                suffix_pos=lastpos+torch.arange(1,len(suffix)+1,device=ids.device).view(1,1,-1)
                if execution=='full' and replay_tokens is not None:
                    extra=torch.tensor([generated+suffix],device=ids.device)
                    all_embeds=torch.cat((embeds,self.model.get_input_embeddings()(extra)),dim=1)
                    extra_pos=pos[...,-1:]+torch.arange(1,extra.shape[1]+1,device=ids.device).view(1,1,-1)
                    h=self.lm(inputs_embeds=all_embeds,position_ids=torch.cat((pos,extra_pos),dim=-1),use_cache=False).last_hidden_state
                else:
                    h=self.lm(input_ids=suffix_ids,position_ids=suffix_pos,past_key_values=cache,use_cache=True).last_hidden_state
                logits=self.model.lm_head(h[:,-1]).float()[0]
                letter_ids=[]
                for c in 'ABCDEFGHIJ'[:n_letters]:
                    token_ids=self.t.encode(c,add_special_tokens=False)
                    if len(token_ids)!=1:raise ValueError('non-single letter token')
                    letter_ids.append(token_ids[0])
                chosen=logits[letter_ids];probs=chosen.softmax(-1)
                return {'cap_checkpoints':snapshots,'letter_logits':chosen.cpu().tolist(),'probabilities':probs.cpu().tolist(),'mass':float(logits.softmax(-1)[letter_ids].sum()),'prediction':int(probs.argmax()),'generated_tokens':generated,'generation_tokens':len(generated),'text':self.t.decode(generated),'finish_reason':stop,'forced_answer':True,'execution':execution,'prompt_tokens':int(ids.shape[1]),'input_tokens':int(ids.shape[1])-p+len(suffix),'branch_tokens':int(ids.shape[1])-p+len(suffix),'reused_prefix_tokens':p if execution=='cached' else 0,'adapter_enabled':False,'identity':self.reasoning_identity()}
        finally:self.ctx['start']=previous

    def exit_features(self,state,block,n_letters,layers=(31,47,55)):
        """Full-depth feature capture only; explicitly does not execute early exit."""
        captured={};handles=[]
        def hook(layer):
            def capture(module,args,output):
                value=output[0] if isinstance(output,tuple) else output
                captured[layer]=value[:,-1].detach().clone()
            return capture
        try:
            for layer in layers:handles.append(self.lm.layers[layer].register_forward_hook(hook(layer)))
            full=self.ask(state,block,n_letters)
            result={}
            with self.torch.inference_mode():
                text=self._render(state['parts'],block);raw=self.t.encode(text,add_special_tokens=False)
                letter_ids=[self.t.encode(text+c,add_special_tokens=False)[-1] for c in 'ABCDEFGHIJ'[:n_letters]]
                for layer,value in captured.items():
                    logits=self.model.lm_head(self.lm.norm(value)).float()[0,letter_ids]
                    result[str(layer)]={'feature':value.float().cpu().tolist()[0],'letter_logits':logits.cpu().tolist(),'prediction':int(logits.argmax())}
            return {'full':{k:(v.tolist() if hasattr(v,'tolist') else v) for k,v in full.items()},'layers':result,'executed_truncated':False,'total_layers':len(self.lm.layers)}
        finally:
            for handle in handles:handle.remove()