LFM-MD-V1-VL-3B / physical_source.py
win10's picture
Release merged checkpoint 110 with persistent memory runtime and detailed model card
51bb18e verified
Raw History Blame Contribute Delete
19.6 kB
"""Source-only FFN weight frames with independently verified reconstruction.
The frozen physical prior is shared. Each observed frame owns a small residual
adapter; only positions, lengths and digests accompany those weights. Readers
never supply observed tokens, VAE codes or saved KV states to reconstruction.
"""
from dataclasses import dataclass,replace
import hashlib,json,math
import torch
from torch.nn.utils.rnn import pad_sequence
@dataclass(frozen=True)
class PhysicalSource:
count: int
checksum: str
def __post_init__(self):
if type(self.count) is not int or self.count<1:raise ValueError('positive physical source length required')
if not isinstance(self.checksum,str) or len(self.checksum)!=64:raise ValueError('physical source requires a SHA256')
try:raw=bytes.fromhex(self.checksum)
except ValueError:raise ValueError('physical source requires a SHA256') from None
if len(raw)!=32:raise ValueError('physical source requires a SHA256')
object.__setattr__(self,'checksum',raw.hex())
@dataclass(frozen=True)
class PhysicalFrame:
offset: int
count: int
checksum: str
record: int
state: object
def detach(self,device=None):
from .episodic_adapters import AdapterMemory
groups=[{k:v.detach().to(device) if device is not None else v.detach() for k,v in g.items()}
for g in (self.state.factors,self.state.first_moment,self.state.second_moment)]
return replace(self,state=AdapterMemory(*groups,self.state.commits).detach())
@dataclass(frozen=True)
class FrameSpan:
record: int
offset: int
count: int
@dataclass(frozen=True)
class PhysicalRecall:
pieces: tuple # (record, offset, verified token tensor)
frames: tuple # all frame descriptors, including failed reads
@property
def verified_positions(self):return sum(len(ids) for _,_,ids in self.pieces)
@property
def positions(self):return sum(f.count for f in self.frames)
def describe_source(ids):
if ids.ndim!=1 or not len(ids):raise ValueError('one complete observed source required')
raw=ids.detach().cpu().to(torch.int64).contiguous().numpy().tobytes()
return PhysicalSource(len(ids),hashlib.sha256(raw).hexdigest())
def frame_metadata(frames):
return [dict(offset=f.offset,count=f.count,checksum=f.checksum,record=f.record,commits=f.state.commits,
rank=next(v.shape[0] for k,v in f.state.factors.items() if k.endswith('.A'))) for f in frames]
def describe_frames(frames):
if not frames:return None
return PhysicalSource(sum(f.count for f in frames),hashlib.sha256(
json.dumps(frame_metadata(frames),sort_keys=True,separators=(',',':')).encode()).hexdigest())
def source_tensors(source):
if source is None:return {}
return {'physical_source.count':torch.tensor([source.count],dtype=torch.int64),
'physical_source.checksum':torch.tensor(list(bytes.fromhex(source.checksum)),dtype=torch.uint8)}
def load_source(tensors):
if 'physical_source.count' not in tensors:return None
count=tensors['physical_source.count'];checksum=tensors['physical_source.checksum']
if count.dtype!=torch.int64 or count.shape!=(1,) or checksum.dtype!=torch.uint8 or checksum.shape!=(32,):
raise ValueError('invalid physical source descriptor tensors')
return PhysicalSource(int(count[0]),bytes(checksum.tolist()).hex())
def frame_tensors(frames):
return {f'frame.{i}.{kind}.{name}':value.detach().cpu().contiguous()
for i,f in enumerate(frames) for kind,group in [('factor',f.state.factors),('first',f.state.first_moment),('second',f.state.second_moment)]
for name,value in group.items()}
def frame_shapes(model,metadata):
if not isinstance(metadata,list):raise ValueError('physical frame metadata must be a list')
shapes={};offset=0;last_record=-1
rank=model.config.physical_frame_rank;limit=model.config.physical_frame_tokens
for i,f in enumerate(metadata):
if (type(f.get('offset')) is not int or f['offset']!=offset or type(f.get('count')) is not int
or not 0<f['count']<=limit or f.get('rank')!=rank or type(f.get('record')) is not int
or f['record'] not in (last_record,last_record+1) or f['record']<0
or type(f.get('commits')) is not int or f['commits']<1):
raise ValueError('invalid physical frame layout')
PhysicalSource(f['count'],f['checksum']);offset+=f['count'];last_record=f['record']
for name,(inputs,outputs) in model._episodic_adapter_bank.targets.items():
for suffix,shape in [('A',[rank,inputs]),('B',[outputs,rank])]:
for kind in ['factor','first','second']:shapes[f'frame.{i}.{kind}.{name}.{suffix}']=shape
return shapes
def load_frames(model,metadata,tensors):
from .episodic_adapters import AdapterMemory
frames=[]
for i,f in enumerate(metadata):
groups=[{name[len(prefix):]:value for name,value in tensors.items() if name.startswith(prefix)}
for prefix in (f'frame.{i}.factor.',f'frame.{i}.first.',f'frame.{i}.second.')]
if any(v.dtype!=torch.float32 or not torch.isfinite(v).all() for g in groups for v in g.values()):
raise ValueError('invalid physical frame values')
if any((v<0).any() for v in groups[2].values()):raise ValueError('negative frame second moment')
# Keep archived frames on CPU. Only a frame being decoded moves to GPU.
state=AdapterMemory(*groups,f['commits']).detach()
frames.append(PhysicalFrame(f['offset'],f['count'],f['checksum'],f['record'],state))
return tuple(frames)
def effective_adapter(model,state,*,anchor=None):
from .episodic_adapters import AdapterMemory
anchor=model.physical_memory.factors(create_graph=torch.is_grad_enabled()) if anchor is None else anchor
values={k:torch.cat((anchor[k],v.to(anchor[k])),dim=0 if k.endswith('.A') else 1) for k,v in state.factors.items()}
return AdapterMemory(values,{}, {},state.commits)
def initial_frame(model,*,create_graph):
from .episodic_adapters import AdapterMemory
prior=model.physical_memory.factors(create_graph=create_graph);rank=model.config.physical_frame_rank
values={k:(v[:rank].clone() if k.endswith('.A') else v.new_zeros((v.shape[0],rank))).requires_grad_(True)
for k,v in prior.items()}
return AdapterMemory(values,{k:torch.zeros_like(v) for k,v in values.items()},
{k:torch.zeros_like(v) for k,v in values.items()})
def attach_first_order_gradient(model,initial,numerical,directions):
from .episodic_adapters import AdapterMemory
values={}
for name,value in numerical.factors.items():
rate=model.physical_memory.rate(name);start=initial.factors[name]
values[name]=value.detach()+(start-start.detach())-(rate-rate.detach())*directions[name]
return AdapterMemory(values,numerical.first_moment,numerical.second_moment,numerical.commits)
@torch.no_grad()
def update_frame_batch(model,states,gradients,counts,learning_rate,rates=None):
"""Vectorize the same independently clipped Adam updates across frames."""
from .episodic_adapters import AdapterMemory
names=tuple(states[0].factors);width=len(names);step=states[0].commits+1
if any(tuple(s.factors)!=names or s.commits+1!=step for s in states):raise ValueError('frame update batch differs')
if len(gradients)!=len(states)*width:raise ValueError('missing physical frame gradients')
values={name:torch.stack([gradients[i*width+j].float() for i in range(len(states))])*counts.sum()/counts[:,None,None]
for j,name in enumerate(names)}
norm=torch.stack([v.square().sum((1,2)) for v in values.values()]).sum(0).sqrt()
if not torch.isfinite(norm).all():raise FloatingPointError('nonfinite physical frame gradients')
clip=1./norm.clamp_min(1.);factors={};first={};second={};directions={}
for name,g in values.items():
g=g*clip[:,None,None]
first[name]=.9*torch.stack([s.first_moment[name] for s in states])+.1*g
second[name]=.999*torch.stack([s.second_moment[name] for s in states])+.001*g.square()
update=(first[name]/(1-.9**step))/(second[name]/(1-.999**step)).sqrt().add(1e-8)
rate=model.physical_memory.rate(name).detach() if rates is None else rates[name]
factors[name]=torch.stack([s.factors[name] for s in states])-learning_rate*(update*rate)
directions[name]=learning_rate*update
result=[AdapterMemory({k:v[i].detach().requires_grad_(True) for k,v in factors.items()},
{k:v[i] for k,v in first.items()},{k:v[i] for k,v in second.items()},step) for i in range(len(states))]
return result,directions
def frame_sources(model,sources):
bos=getattr(model.config,'bos_token_id',None)
if bos is None:bos=model.config.text_config.bos_token_id
if bos is None:raise ValueError('physical reconstruction requires a fixed model BOS')
framed=[torch.cat((s.new_tensor([bos]),s)) for s in sources]
ids=pad_sequence(framed,batch_first=True,padding_value=0)
mask=torch.arange(ids.shape[1],device=ids.device)[None]<ids.new_tensor(list(map(len,framed)))[:,None]
labels=ids.masked_fill(~mask,-100);labels[:,0]=-100
return ids,mask,labels
def source_forward(model,states,sources,*,create_graph,anchor=None):
from .episodic_adapters import AdapterBatch
from .autonomous_memory import fixed_native_source_forward
ids,mask,labels=frame_sources(model,sources)
if model._memory_context.get() is not None:raise RuntimeError('finish active memory reads before writing frames')
if anchor is None:anchor=model.physical_memory.factors(create_graph=create_graph)
with torch.set_grad_enabled(create_graph),model._episodic_adapter_bank.use(
AdapterBatch(tuple(effective_adapter(model,s,anchor=anchor) for s in states))):
output=fixed_native_source_forward(model,input_ids=ids,attention_mask=mask,logits_to_keep=1,
output_hidden_states=True,use_cache=False,return_dict=True)
return output.hidden_states[-1],labels
def fit_sources(model,native_forward,units,sources,*,steps=64,learning_rate=.001,create_graph=False):
"""Incoming sources only; append independent residual weights without replay.
First-order meta gradients are represented by the initial factors and the
accumulated detached Adam directions. This preserves the existing first-
order derivative without retaining every intermediate optimizer tape.
"""
from .memory_recall_training import source_causal_loss,evaluation_write_context
from .episodic_adapters import AdapterMemory
if type(steps) is not int or steps<1 or not math.isfinite(learning_rate) or learning_rate<=0:
raise ValueError('positive source fitting budget and rate required')
if len(units)!=len(sources):raise ValueError('source/unit batch mismatch')
limit=model.config.physical_frame_tokens;parts=[];owners=[]
for row,source in enumerate(sources):
if source.ndim!=1 or not len(source):raise ValueError('nonempty source required')
for start in range(0,len(source),limit):parts.append(source[start:start+limit]);owners.append((row,start))
initial=[initial_frame(model,create_graph=create_graph) for _ in parts]
states=[s.detach() for s in initial]
directions=None
anchor={k:v.detach() for k,v in model.physical_memory.factors(create_graph=False).items()}
rates={k:model.physical_memory.rate(k).detach() for k in states[0].factors}
losses=[];checkpointed=0;bank=model._episodic_adapter_bank
for _ in range(steps):
with evaluation_write_context(model) as checkpointed,torch.enable_grad():
hidden,labels=source_forward(model,states,parts,create_graph=True,anchor=anchor)
loss=source_causal_loss(model,hidden,labels,fixed_projection=True)
factors=tuple(p for s in states for p in s.factors.values())
gradients=torch.autograd.grad(loss,factors,create_graph=False)
counts=(labels[:,1:]!=-100).sum(1)
states,update=update_frame_batch(model,states,gradients,counts,learning_rate,rates)
if directions is None:directions=update
else:
for name in directions:directions[name].add_(update[name])
losses.append(float(loss.detach()))
result=list(units);added=[[] for _ in units]
for i,((row,offset),part,s,origin) in enumerate(zip(owners,parts,states,initial)):
if create_graph:
s=attach_first_order_gradient(model,origin,s,{k:v[i] for k,v in directions.items()})
prior=units[row].frames;start=sum(f.count for f in prior);record=prior[-1].record+1 if prior else 0
frame=PhysicalFrame(start+offset,len(part),describe_source(part).checksum,record,s)
added[row].append(frame if create_graph else frame.detach('cpu'))
for row,unit in enumerate(units):
frames=unit.frames+tuple(added[row]);result[row]=replace(unit,frames=frames,source=describe_frames(frames))
return result,dict(source_losses=losses,source_updates=steps,checkpointed_layers=checkpointed,
source_bos_supervised=True,query_or_answer_seen=False,physical_frames=len(parts),
physical_frame_rank=model.config.physical_frame_rank,physical_frame_tokens=limit,
inner_loop='source-only first-order Adam; frozen native body and inherited prior')
def teacher_source_view(model,native_forward,units,sources,*,create_graph):
from torch.utils.checkpoint import checkpoint
from .memory_recall_training import source_causal_loss
from .sft_lora import checkpoint_weight_contexts
from .autonomous_memory import fixed_output_projection
parts=[];states=[];owners=[]
for row,(unit,source) in enumerate(zip(units,sources)):
record=unit.frames[-1].record
selected=[f for f in unit.frames if f.record==record]
if sum(f.count for f in selected)!=len(source):raise ValueError('teacher source does not match the latest write')
offset=0
for frame in selected:
part=source[offset:offset+frame.count];offset+=frame.count
if describe_source(part).checksum!=frame.checksum:raise ValueError('teacher frame identity differs')
parts.append(part);states.append(frame.state);owners.append(row)
with torch.set_grad_enabled(create_graph):
hidden,labels=source_forward(model,states,parts,create_graph=create_graph)
loss=source_causal_loss(model,hidden,labels,fixed_projection=True)
dictionary=model.get_input_embeddings().weight.detach();prefixes=[[] for _ in units]
def decode(x):
logits=fixed_output_projection(model,x).float()
with torch.autocast(device_type=x.device.type,enabled=False):return logits.softmax(-1).to(dictionary.dtype)@dictionary
for i,(row,part) in enumerate(zip(owners,parts)):
x=hidden[i,:len(part)]
prefixes[row].append(checkpoint(decode,x,use_reentrant=False,context_fn=checkpoint_weight_contexts) if create_graph else decode(x))
return [torch.cat(p) for p in prefixes],loss
@torch.no_grad()
def reconstruct(model,unit):
from .sft_lora import shared_effective_weights
if not unit.frames or describe_frames(unit.frames)!=unit.source:raise ValueError('physical frame manifest mismatch')
bos=getattr(model.config,'bos_token_id',None)
if bos is None:bos=model.config.text_config.bos_token_id
ids=torch.tensor([[bos]],device=model.device);pieces=[];receipts=[]
archive=model._archive_context.set(None);decoding=model._physical_source_context.set(True)
try:
with shared_effective_weights():
anchor=model.physical_memory.factors(create_graph=False)
for frame in unit.frames:
with model._episodic_adapter_bank.use(effective_adapter(model,frame.state,anchor=anchor)):
output=model.generate(input_ids=ids,attention_mask=torch.ones_like(ids),
max_new_tokens=frame.count,eos_token_id=None,do_sample=False,use_cache=True)
restored=output[0,1:];verified=len(restored)==frame.count and describe_source(restored).checksum==frame.checksum
if verified:pieces.append((frame.record,frame.offset,restored))
receipts.append(dict(record=frame.record,offset=frame.offset,count=frame.count,verified=verified))
finally:model._physical_source_context.reset(decoding);model._archive_context.reset(archive)
recall=PhysicalRecall(tuple(pieces),tuple(FrameSpan(f.record,f.offset,f.count) for f in unit.frames))
return recall,dict(verified=recall.verified_positions==recall.positions,verified_frames=sum(r['verified'] for r in receipts),
total_frames=len(receipts),verified_positions=recall.verified_positions,source_positions=recall.positions,
frames=receipts,persistent_port_slots_read=False,persistent_source_or_kv_read=False,fixed_bos=True,
decoder_weights='independent FFN residuals anchored to the inherited physical prior')
def recalled_tokens(model,recalls):
"""Keep every verified fragment and mark gaps; never join across a missing span."""
chunks=[];device=model.device
gap=torch.tensor(model.config.physical_gap_token_ids,device=device,dtype=torch.long)
separator=torch.tensor(model.config.physical_record_separator_ids,device=device,dtype=torch.long)
for recall in recalls:
by_record={}
for record,offset,ids in recall.pieces:by_record.setdefault(record,[]).append((offset,ids))
for record in sorted({f.record for f in recall.frames}):
frames=[f for f in recall.frames if f.record==record];parts=by_record.get(record,[])
if not parts:continue
if chunks:chunks.append(separator)
expected=frames[0].offset
for offset,ids in sorted(parts,key=lambda p:p[0]):
if offset!=expected:chunks.append(gap)
chunks.append(ids.to(device));expected=offset+len(ids)
if expected!=frames[-1].offset+frames[-1].count:chunks.append(gap)
return torch.cat(chunks) if chunks else torch.empty(0,device=device,dtype=torch.long)
def prepare_physical_query(model,restored,inputs,*,max_new_tokens=None,max_length=None):
from .sequence_memory import assemble_memory_query,MemoryContextLimitError
if inputs.get('pixel_values') is not None:
from .sequence_memory import native_query_embeddings
query=native_query_embeddings(model,inputs,memory_state=None,port_memory_state=None,use_memory=False)
else:query=model.get_input_embeddings()(inputs['input_ids'])[0]
prefix=model.get_input_embeddings()(restored)
bos=getattr(model.config,'bos_token_id',None)
if bos is None:bos=getattr(model.config.text_config,'bos_token_id',None)
value,_=assemble_memory_query(prefix,query,inputs['input_ids'][0],bos)
reserve=max_new_tokens;config=inputs.get('generation_config') or model.generation_config
if reserve is None:reserve=config.max_new_tokens
if reserve is None:reserve=max(0,(max_length or config.max_length)-inputs['input_ids'].shape[1])
if len(value)+reserve>model.config.text_config.max_position_embeddings:
raise MemoryContextLimitError('verified physical fragments plus query/output exceed native context')
result=dict(inputs,inputs_embeds=value[None],attention_mask=torch.ones((1,len(value)),device=value.device,dtype=torch.long))
for key in ('pixel_values','spatial_shapes','pixel_attention_mask'):result.pop(key,None)
return result