LFM-MD-V1-VL-3B / sequence_memory.py
win10's picture
Release merged checkpoint 110 with persistent memory runtime and detailed model card
51bb18e verified
Raw History Blame Contribute Delete
10.1 kB
"""Ordered native-feature slots; no VAE, fitted decoder, or numerical patches.
Features are explicit memory payload and charged at their actual dtype/size.
These slots preserve native vision embeddings as well as text embeddings.
Physical FFN frames and MaleCNS fast weights remain independent substrates.
"""
from contextlib import contextmanager
from contextvars import ContextVar
from dataclasses import dataclass
import hashlib
import re
import torch
DTYPES={'bfloat16':torch.bfloat16,'float16':torch.float16,'float32':torch.float32}
NATIVE_QUERY_INPUTS=('input_ids','attention_mask','pixel_values','pixel_attention_mask','spatial_shapes')
class NativeInputCaptured(Exception):
"""Caller-local early return after native multimodal input assembly."""
class MemoryContextLimitError(ValueError):
"""Full ordered read does not fit the requested context, without truncation."""
def feature_checksum(x):
return hashlib.sha256(x.detach().cpu().contiguous().view(torch.uint8).numpy().tobytes()).hexdigest()
@dataclass(frozen=True)
class SequenceMemory:
features: torch.Tensor
checksum: str
def tensors(self):return (self.features,)
def detach(self):return SequenceMemory(self.features.detach(),self.checksum)
@property
def count(self):return len(self.features)
@property
def bytes(self):return self.features.numel()*self.features.element_size()
def validate_sequence(model,segment):
if not isinstance(segment,SequenceMemory):raise ValueError('invalid ordered feature memory')
x=segment.features
if (x.ndim!=2 or not len(x) or x.shape[1]!=model.config.text_config.hidden_size
or str(x.dtype).removeprefix('torch.') not in DTYPES or not torch.isfinite(x).all()):
raise ValueError('invalid ordered native features')
if not re.fullmatch('[a-f0-9]{64}',segment.checksum):raise ValueError('invalid feature checksum')
@torch.no_grad()
def encode_sequence(model,features,*,chunk_size=128):
if type(chunk_size) is not int or chunk_size<1:raise ValueError('positive storage chunk size required')
result=[]
for start in range(0,len(features),chunk_size):
value=features[start:start+chunk_size].detach().cpu().contiguous().clone()
segment=SequenceMemory(value,feature_checksum(value));validate_sequence(model,segment);result.append(segment)
if not result:raise ValueError('nonempty observed features required')
return tuple(result)
@torch.no_grad()
def decode_segment(model,segment):
validate_sequence(model,segment)
if feature_checksum(segment.features)!=segment.checksum:raise ValueError('ordered feature checksum mismatch')
return segment.features.detach()
def install_sequence_capture(model):
"""A caller-local capture scope, safe across independent read sessions."""
model._sequence_capture_context = ContextVar(f'sequence_capture_{id(model)}', default=None)
def capture(module, args, kwargs):
scope = model._sequence_capture_context.get()
if scope is None:return
captured, attention_mask, stop_before_language = scope
value = kwargs.get('inputs_embeds')
if value is None:raise ValueError('native language input embeddings were not supplied')
if value.shape[0] != 1:raise ValueError('write one ordered session per call')
value = value[0]
if attention_mask is not None:
if attention_mask.shape != (1, len(value)):raise ValueError('sequence validity mask mismatch')
value = value[attention_mask[0].bool()]
captured.append(value.detach().clone())
if stop_before_language:raise NativeInputCaptured
model.model.language_model.register_forward_pre_hook(capture, with_kwargs=True)
@contextmanager
def capture_native_inputs(model, attention_mask=None, *, stop_before_language=False):
"""Capture native text/vision embeddings without changing their forward."""
if model._sequence_capture_context.get() is not None:raise RuntimeError('nested sequence capture')
captured = []
token = model._sequence_capture_context.set((captured, attention_mask, stop_before_language))
try:yield captured
finally:model._sequence_capture_context.reset(token)
@torch.no_grad()
def native_query_embeddings(model,inputs,*,memory_state,port_memory_state,use_memory=True):
"""Use native vision/projector/merge code, stopping before language layers.
The permanent capture hook uses ContextVar state. No hooks are added or
removed on a shared module while other callers may be using it.
"""
from .dnc_memory import visual_slots
query={k:v for k,v in inputs.items() if k in NATIVE_QUERY_INPUTS}
with capture_native_inputs(model,stop_before_language=True) as captured:
try:
model(**query,memory_state=memory_state,
port_memory_state=visual_slots(model,port_memory_state) if use_memory else None,
use_memory=use_memory,use_cache=False,logits_to_keep=1)
except NativeInputCaptured:pass
if len(captured)!=1:raise RuntimeError('native query input assembly was not captured')
return captured[0]
@torch.no_grad()
def prepare_ordered_inputs(model, sequences, inputs, *, max_new_tokens=None, max_length=None,
query_embedding_provider=None):
"""Restore selected sequence features in order through ALL native layers.
No prefix token IDs are recovered. Original input_ids remain the returned
generation prefix, so existing callers keep their normal output slicing.
Recurrent and attention caches are rebuilt normally and remain ephemeral.
"""
if not sequences:raise ValueError('this unit has no ordered memory; re-observe its source to create it')
forbidden = {'inputs_embeds', 'past_key_values', 'position_ids'}
if any(inputs.get(k) is not None for k in forbidden):
raise ValueError('ordered recall requires fresh native query inputs')
if inputs.get('pixel_values') is not None and query_embedding_provider is None:
raise ValueError('a visual query requires native query input assembly')
ids = inputs.get('input_ids')
if ids is None or ids.ndim != 2 or ids.shape[0] != 1:raise ValueError('ordered recall requires one tokenized query')
mask = inputs.get('attention_mask')
if mask is None:mask = torch.ones_like(ids)
if mask.shape != ids.shape or not mask.bool().all():raise ValueError('ordered recall requires an unpadded query')
length = sum(s.count for s in sequences)
generation_config = inputs.get('generation_config') or model.generation_config
reserve = max_new_tokens if max_new_tokens is not None else generation_config.max_new_tokens
if reserve is None:
effective_max = max_length if max_length is not None else generation_config.max_length
reserve = max(0, effective_max - ids.shape[1])
limit = model.config.text_config.max_position_embeddings
if ids.shape[1]+reserve>limit:
raise ValueError(f'query + requested output ({ids.shape[1]+reserve}) exceeds native context {limit}; no truncation performed')
if length + ids.shape[1] + reserve > limit:
raise MemoryContextLimitError(f'selected memory + query + output ({length + ids.shape[1] + reserve}) exceeds native context {limit}; no truncation performed')
embeddings = model.get_input_embeddings()(ids) if query_embedding_provider is None else query_embedding_provider()[None]
if embeddings.shape!=(1,ids.shape[1],model.config.text_config.hidden_size):
raise ValueError('native query embedding shape differs from its token sequence')
prefix = torch.cat([decode_segment(model, segment) for segment in sequences]).to(embeddings)
bos=getattr(model.config,'bos_token_id',None)
if bos is None:bos=getattr(model.config.text_config,'bos_token_id',None)
combined,_=assemble_memory_query(prefix,embeddings[0],ids[0],bos)
# The query image has already traversed native vision and image merging.
# Passing its pixels again would merge it at unshifted prefix positions.
prepared={k:v for k,v in inputs.items() if k not in ('pixel_values','pixel_attention_mask','spatial_shapes')}
return dict(prepared, inputs_embeds=combined[None],
attention_mask=torch.ones((1, length + ids.shape[1]), device=ids.device, dtype=mask.dtype))
def assemble_memory_query(prefix,query,query_ids,bos_token_id):
"""Keep the native start token before recalled observations and the query.
No feature or query token is omitted. Inserting a new conversation start
after the recalled observations can make the model treat them as outside
the current conversation. The same ordering is used by training and recall.
"""
leading=int(bos_token_id is not None and len(query_ids)>0 and int(query_ids[0])==bos_token_id)
return torch.cat((query[:leading],prefix,query[leading:]),dim=0),leading
def sequence_tensors(segments):
return {f'sequence.{i}.features':s.features.detach().cpu().contiguous() for i,s in enumerate(segments)}
def sequence_metadata(segments):
return [dict(count=s.count,dtype=str(s.features.dtype).removeprefix('torch.'),checksum=s.checksum) for s in segments]
def sequence_shapes(model,metadata):
if not isinstance(metadata,list):raise ValueError('invalid ordered feature metadata')
result={}
for i,row in enumerate(metadata):
n=row.get('count')
if type(n) is not int or n<1 or row.get('dtype') not in DTYPES:
raise ValueError('invalid ordered feature shape metadata')
result[f'sequence.{i}.features']=[n,model.config.text_config.hidden_size]
return result
def load_sequences(model,metadata,tensors):
result=[]
sequence_shapes(model,metadata)
for i,row in enumerate(metadata):
value=tensors[f'sequence.{i}.features']
if value.dtype!=DTYPES[row['dtype']]:raise ValueError('ordered feature dtype differs from manifest')
segment=SequenceMemory(value,row['checksum']);decode_segment(model,segment);result.append(segment)
return tuple(result)