File size: 12,700 Bytes
105f9ef
 
d144fee
 
 
 
 
 
 
 
 
 
105f9ef
d144fee
105f9ef
 
d144fee
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
105f9ef
d144fee
105f9ef
 
 
 
 
d144fee
 
 
 
 
105f9ef
d144fee
 
 
105f9ef
 
 
 
 
d144fee
 
 
105f9ef
 
 
 
 
 
 
 
 
 
 
 
 
 
 
d144fee
 
 
 
 
105f9ef
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
d144fee
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
105f9ef
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
"""CUDA contract-v3 engine. Separate identity; never presents itself as MLX.
Model dependencies import lazily. Mutable cache branches are always deep-copied.

v1.1 (Solomon v1.1): an explicit precision, recorded in the identity.
  fp32  float32 weights and compute (the qualified numerics are installed on top by solomon/engine_numerics.py)
  bf16  bfloat16 weights and compute; the gated-delta recurrence is promoted to float32 (mode 'bf16-fp32-recurrence',
        the same mode the Solomon v1.1 trainer runs in). DEFAULT for v1.1 serving.
  int8  bf16 as above, plus torchao Int8WeightOnly on every nn.Linear inside language_model.layers (vision tower,
        embeddings and lm_head stay bf16); the LoRA adapter wraps the quantised linears unmerged in float32.
The fp32 recurrence promotion is installed once per process, idempotently, and the wrapper carries __wrapped__ so
that inspect.unwrap reaches the reference implementation (re-deriving the promoted wrapper from its own source was
the NameError 'original' crash).
"""
import copy,functools,hashlib,inspect,json
from pathlib import Path

PRECISIONS={'fp32':'fp32','bf16':'bf16-fp32-recurrence','int8':'int8-bf16-fp32-recurrence'}
MODES={v:k for k,v in PRECISIONS.items()}
RECURRENCE=('torch_chunk_gated_delta_rule','torch_recurrent_gated_delta_rule')

def recurrence_kernel(function):
    """'fla' / 'hub-kernel' / 'torch-reference': the implementation transformers dispatches this function to."""
    target=function
    while getattr(target,'_solomon_promoted',False):target=target.__wrapped__
    free=getattr(target,'__code__',None)
    cells=dict(zip(free.co_freevars,[c.cell_contents for c in (target.__closure__ or ())])) if free is not None else {}
    implementation=cells.get('implementation')
    new=cells['is_new_implementation'] if 'is_new_implementation' in cells else (implementation is not None and implementation is not getattr(target,'__wrapped__',None))
    if not new:return 'torch-reference'
    return 'fla' if 'fla' in (getattr(implementation,'__module__','') or '') else 'hub-kernel'

def promote_recurrence(modeling,torch):
    """Install the fp32 recurrence promotion (idempotent). Returns the dispatched kernel name."""
    for name in RECURRENCE:
        original=getattr(modeling,name)
        if getattr(original,'_solomon_promoted',False):continue
        def promoted(q,k,v,*args,_original=original,**kwargs):
            old=q.dtype
            for key in ('g','beta','initial_state'):
                if isinstance(kwargs.get(key),torch.Tensor):kwargs[key]=kwargs[key].float()
            out,state=_original(q.float(),k.float(),v.float(),*args,**kwargs)
            return out.to(old),state
        functools.update_wrapper(promoted,original);promoted._solomon_promoted=True
        setattr(modeling,name,promoted)
    return recurrence_kernel(getattr(modeling,RECURRENCE[0]))

def unpromote_recurrence(modeling):
    """Restore the functions promote_recurrence replaced (an fp32 engine after a bf16 one in the same process)."""
    for name in RECURRENCE:
        f=getattr(modeling,name)
        while getattr(f,'_solomon_promoted',False):f=f.__wrapped__
        setattr(modeling,name,f)

def int8_config():
    from torchao.quantization import quantize_
    try:
        from torchao.quantization import Int8WeightOnlyConfig;return quantize_,Int8WeightOnlyConfig()
    except ImportError:
        from torchao.quantization import int8_weight_only;return quantize_,int8_weight_only()

def load_int8(Auto,model_dir,torch):
    """bf16 load on CPU, torchao Int8WeightOnly per decoder layer on the GPU, then move (fits a 48 GB card)."""
    from torch import nn
    quantize_,config=int8_config()
    model=Auto.from_pretrained(model_dir,dtype=torch.bfloat16,device_map='cpu',attn_implementation='sdpa')
    count=0
    for layer in model.model.language_model.layers:
        layer.to('cuda');count+=sum(isinstance(m,nn.Linear) for m in layer.modules())
        quantize_(layer,config,filter_fn=lambda m,fqn:isinstance(m,nn.Linear));torch.cuda.empty_cache()
    return model.to('cuda').eval(),count

class CudaEngine:
    def __init__(self,model_dir='base',adapter='adapter/adapter.safetensors',mode='fp32',placement='question',precision=None):
        import torch
        from torch import nn
        from transformers import AutoModelForImageTextToText,AutoProcessor
        from transformers.cache_utils import LinearAttentionLayer
        from safetensors.torch import load_file
        if precision is not None:
            if precision not in PRECISIONS:raise ValueError('precision must be one of '+', '.join(PRECISIONS))
            mode=PRECISIONS[precision]
        if mode not in MODES:raise ValueError('unknown arithmetic mode '+str(mode))
        self.torch=torch;self.mode=mode;self.precision=MODES[mode];self.ctx={'start':None};self.adapter=adapter;self.placement=placement
        torch.backends.cuda.matmul.allow_tf32=False
        quantised=None
        if self.precision=='int8':self.model,quantised=load_int8(AutoModelForImageTextToText,model_dir,torch)
        else:self.model=AutoModelForImageTextToText.from_pretrained(model_dir,dtype=torch.float32 if mode=='fp32' else torch.bfloat16,device_map='cuda',attn_implementation='sdpa').eval()
        self.processor=AutoProcessor.from_pretrained(model_dir);self.t=self.processor.tokenizer;self.lm=self.model.model.language_model
        def update(cache,recurrent_states,state_idx=0,**kwargs):
            if not cache.is_recurrent_states_initialized[state_idx]:cache.lazy_initialization(recurrent_states=recurrent_states,state_idx=state_idx)
            cache.recurrent_states[state_idx]=recurrent_states;return recurrent_states
        LinearAttentionLayer.update_recurrent_state=update
        from transformers.models.qwen3_5 import modeling_qwen3_5 as modeling
        if mode=='fp32':unpromote_recurrence(modeling);kernel=recurrence_kernel(getattr(modeling,RECURRENCE[0]))
        else:kernel=promote_recurrence(modeling,torch)
        if adapter:
            ctx=self.ctx;w=load_file(adapter)
            class LoRA(nn.Module):
                def __init__(self,linear,a,b):super().__init__();self.linear=linear;self.a=a.cuda().float();self.b=b.cuda().float()
                def forward(self,x):
                    y=self.linear(x);s=ctx['start']
                    if s is None or s>=x.shape[1]:return y
                    z=(2*((x[:,s:].float()@self.a)@self.b)).to(y.dtype)
                    return y+z if s==0 else torch.cat((y[:,:s],y[:,s:]+z),dim=1)
            for name in sorted({k.rsplit('.',1)[0] for k in w}):
                parts=name.split('.');owner=self.lm.layers[int(parts[2])]
                for p in parts[3:-1]:owner=getattr(owner,p)
                setattr(owner,parts[-1],LoRA(getattr(owner,parts[-1]),w[name+'.lora_a'],w[name+'.lora_b']))
        sha=lambda p:hashlib.sha256(Path(p).read_bytes()).hexdigest()
        self.identity={'backend':'cuda','execution':'cached','arithmetic':mode,'placement':placement,'model_sha256':sha(Path(model_dir)/'model.safetensors.index.json'),'adapter_sha256':sha(adapter) if adapter else None,'code_sha256':sha(__file__),'torch':torch.__version__}
        # Precision identity (v1.1). Only added for non-fp32 precisions, so an fp32 identity keeps the v1 key set.
        if self.precision!='fp32':
            self.identity.update(precision=self.precision,weights='int8-weight-only-per-channel(torchao) language_model.layers nn.Linear; rest bf16' if quantised is not None else 'bf16',
                                 numerics='bf16-sdpa-attention-fp32-recurrence-'+kernel)
            if quantised is not None:self.identity['int8_linears']=int(quantised)
        self.identity['fingerprint']=hashlib.sha256(json.dumps(self.identity,sort_keys=True).encode()).hexdigest()

    def _render(self,parts,block):
        from solomon.engine_contract import SYSTEM,PAGE
        content=''
        for i,p in enumerate(parts):
            if 'text' in p:content+= ('\n' if i and 'image' in parts[i-1] else '')+p['text']
            else:content+= ('\n' if i and 'text' in parts[i-1] else '')+PAGE
        return self.t.apply_chat_template([{'role':'system','content':SYSTEM},{'role':'user','content':'Document:\n'+content+'\n\n'+block}],tokenize=False,add_generation_prompt=True,enable_thinking=False)

    def _encode(self,parts,text,features=None):
        from PIL import Image
        torch=self.torch
        images=[Image.open(p['image']).convert('RGB') for p in parts if 'image' in p]
        enc=self.processor(text=[text],images=images or None,return_tensors='pt',add_special_tokens=False).to('cuda')
        ids=enc['input_ids'];embeds=self.model.get_input_embeddings()(ids)
        if images:
            if features is None:features=torch.cat(self.model.model.get_image_features(enc['pixel_values'],enc['image_grid_thw'],return_dict=True).pooler_output,dim=0)
            features=features.to(embeds);mask,_=self.model.model.get_placeholder_mask(ids,inputs_embeds=embeds,image_features=features)
            embeds=embeds.masked_scatter(mask,features)
        pos,_=self.model.model.get_rope_index(ids,enc.get('mm_token_type_ids',torch.zeros_like(ids)),image_grid_thw=enc.get('image_grid_thw'),attention_mask=enc.get('attention_mask'))
        return ids,embeds,pos,features

    def prefill(self,document):
        from solomon.engine_contract import as_parts
        from transformers import DynamicCache
        parts=as_parts(document);text=self._render(parts,'X');end=text.rfind('\n\nX')
        raw=self.t.encode(text[:end],add_special_tokens=False)
        if 'text' in parts[-1]:raw=raw[:-1]
        with self.torch.inference_mode():
            ids,embeds,pos,features=self._encode(parts,text)
            p=len(raw)+ids.shape[1]-len(self.t.encode(text,add_special_tokens=False))
            cache=DynamicCache(config=self.lm.config);self.ctx['start']=0 if self.adapter and self.placement=='full' else None
            # v1.1: the evidence head reads the document prefix states (final layer, and any tapped layer through
            # the frozen final norm, as tap_layers does for the branch). Kept on the device with the warm state.
            mids={};handles=[]
            def tap(index):
                def capture(module,args,output):
                    h=output[0] if isinstance(output,tuple) else output
                    mids[str(index)]=self.lm.norm(h)[0].detach()
                return capture
            try:
                for index in getattr(self,'prefix_layers',()):handles.append(self.lm.layers[index].register_forward_hook(tap(index)))
                out=self.lm(inputs_embeds=embeds[:,:p],position_ids=pos[...,:p],past_key_values=cache,use_cache=True)
            finally:
                for handle in handles:handle.remove()
        state={'parts':parts,'prefix_ids':ids[:,:p].clone(),'cache':cache,'features':features,'prefix_tokens':p}
        if getattr(self,'keep_prefix_hidden',False):state.update(doc_hidden=out.last_hidden_state[0].detach(),doc_layers=mids,prompt=text)
        return state

    def ask(self,state,block,n_letters,execution='cached'):
        if execution not in ('cached','full'):raise ValueError('invalid execution')
        torch=self.torch;text=self._render(state['parts'],block);raw=self.t.encode(text,add_special_tokens=False)
        letters=[]
        for c in 'ABCDEFGHIJ'[:n_letters]:
            ext=self.t.encode(text+c,add_special_tokens=False)
            if ext[:-1]!=raw or len(ext)!=len(raw)+1:raise ValueError('unstable answer continuation')
            letters.append(ext[-1])
        with torch.inference_mode():
            ids,embeds,pos,_=self._encode(state['parts'],text,state['features']);p=state['prefix_tokens']
            if not torch.equal(ids[:,:p],state['prefix_ids']):raise ValueError('prefix mismatch')
            self.ctx['start']=(0 if execution=='cached' or self.placement=='full' else p) if self.adapter else None
            if execution=='cached':hidden=self.lm(inputs_embeds=embeds[:,p:],position_ids=pos[...,p:],past_key_values=copy.deepcopy(state['cache']),use_cache=True).last_hidden_state
            else:hidden=self.lm(inputs_embeds=embeds,position_ids=pos,use_cache=False).last_hidden_state
            logits=self.model.lm_head(hidden[:,-1]).float()[0];choice=logits[letters];probs=choice.softmax(-1)
        return {'letter_logits':choice.cpu().numpy(),'probabilities':probs.cpu().numpy(),'mass':float(logits.softmax(-1)[letters].sum()),'top_is_letter':int(logits.argmax()) in letters,'execution':execution,'fallback':'','prompt_tokens':int(ids.shape[1]),'branch_tokens':int(ids.shape[1])-p,'reused_prefix_tokens':p if execution=='cached' else 0}