"""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)