Files
cri/export_onnx.py
Tommi Niemi 3c77498dd4 Reframe as Conditioned Reflex Injection (CRI) with multi-model test results
- Rename from "Episodic Memory" to "Conditioned Reflex Injection" throughout
- Make code model-agnostic: --model flag for any HuggingFace backbone
- Support multimodal models (Gemma 4) via .model.language_model resolution
- Add negative stimulus-specificity tests and abstract query tests
- Paper now backed by measured data from 4 backbones:
  Qwen 2.5 0.5B, Gemma 4 E2B-it, E4B-it, E4B base
- Quantization tolerance tested at f32/f16/bf16/int8/int4
- Key findings: smaller base models outperform larger instruct models,
  instruct tuning compresses activation space (hurts discrimination),
  int4 viable if same-precision conditioning/triggering
- Add privacy-by-representation section
- Add Pavlov/Skinner references for conditioning framing

Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
2026-04-06 19:03:32 +07:00

224 lines
7.6 KiB
Python

#!/usr/bin/env python3
"""Export any HuggingFace causal LM as ONNX for epimem.
Exports two models:
1. backbone.onnx: token_ids → hidden_state at target layer (for trigger key computation)
2. lm_head.onnx: full_hidden → logits (for generation)
Model-agnostic: works with Qwen, Gemma, Llama, Mistral, Phi, etc.
For multimodal models (Gemma 4), exports only the text decoder.
Usage:
python export_onnx.py # default: Qwen 2.5 0.5B
python export_onnx.py --model google/gemma-4-E4B-it # Gemma 4
python export_onnx.py --model google/gemma-4-E2B-it # Gemma 4 small
python export_onnx.py --model meta-llama/Llama-3.2-1B # Llama
"""
import argparse
import torch
import torch.nn as nn
import os
DEFAULT_MODEL = "Qwen/Qwen2.5-0.5B"
def resolve_text_model(model):
"""Extract the text decoder from multimodal models, or return as-is.
Gemma 4: ConditionalGeneration → .model (Gemma4Model) → .language_model
LLaVA: ConditionalGeneration → .language_model
Pure CausalLM: model itself
"""
inner = getattr(model, 'model', model)
if hasattr(inner, 'language_model'):
return inner.language_model
if hasattr(model, 'language_model'):
return model.language_model
return model
def get_lm_head_and_norm(model):
"""Find lm_head and final layer norm regardless of architecture.
Returns (norm, lm_head) where norm may be None if not found separately.
"""
text_model = resolve_text_model(model)
# Find lm_head
lm_head = None
if hasattr(text_model, 'lm_head'):
lm_head = text_model.lm_head
elif hasattr(model, 'lm_head'):
lm_head = model.lm_head
# Find final norm — varies by architecture
norm = None
inner = getattr(text_model, 'model', text_model)
for attr in ['norm', 'final_layernorm', 'ln_f', 'final_norm']:
if hasattr(inner, attr):
norm = getattr(inner, attr)
break
return norm, lm_head
class GenericBackboneWrapper(nn.Module):
"""Architecture-agnostic wrapper that extracts hidden state at target layer.
Uses output_hidden_states=True to get intermediate representations
without manually walking architecture-specific layer structures.
Returns:
hidden: hidden state at target_layer (for trigger key)
full_hidden: final layer hidden state (for lm_head / logit computation)
"""
def __init__(self, text_model, target_layer):
super().__init__()
self.text_model = text_model
self.target_layer = target_layer
def forward(self, input_ids):
outputs = self.text_model(input_ids, output_hidden_states=True)
hidden_states = outputs.hidden_states # tuple of (n_layers+1) tensors
target_hidden = hidden_states[self.target_layer] # [B, T, D]
final_hidden = hidden_states[-1] # [B, T, D]
return target_hidden, final_hidden
class GenericLMHead(nn.Module):
"""Architecture-agnostic lm_head wrapper: hidden → logits."""
def __init__(self, norm, lm_head):
super().__init__()
self.norm = norm
self.lm_head = lm_head
def forward(self, full_hidden):
x = full_hidden
if self.norm is not None:
x = self.norm(x)
return self.lm_head(x)
def export(model_name: str, output_dir: str):
from transformers import AutoModelForCausalLM, AutoTokenizer
print(f"Loading {model_name}...")
# Try CausalLM first, fall back to auto for multimodal
is_multimodal = False
try:
model = AutoModelForCausalLM.from_pretrained(
model_name, torch_dtype=torch.float32, trust_remote_code=True)
except (ValueError, KeyError):
from transformers import AutoModel
model = AutoModel.from_pretrained(
model_name, torch_dtype=torch.float32, trust_remote_code=True)
is_multimodal = True
tokenizer = AutoTokenizer.from_pretrained(model_name, trust_remote_code=True)
model.eval()
text_model = resolve_text_model(model)
config = text_model.config
D = config.hidden_size
n_layers = config.num_hidden_layers
V = config.vocab_size
target_layer = n_layers - 1
print(f" arch={type(model).__name__}")
print(f" D={D}, layers={n_layers}, vocab={V}, target_layer={target_layer}")
os.makedirs(output_dir, exist_ok=True)
# ─── Export backbone: tokens → hidden at target layer ─────────
print("Exporting backbone (tokens → hidden)...")
backbone = GenericBackboneWrapper(text_model, target_layer)
dummy_ids = torch.randint(0, min(V, 10000), (1, 32))
with torch.no_grad():
hidden, full = backbone(dummy_ids)
print(f" hidden: {hidden.shape}, full: {full.shape}")
backbone_path = os.path.join(output_dir, "backbone.onnx")
torch.onnx.export(
backbone, dummy_ids,
backbone_path,
input_names=["input_ids"],
output_names=["hidden", "full_hidden"],
dynamic_axes={
"input_ids": {0: "batch", 1: "seq_len"},
"hidden": {0: "batch", 1: "seq_len"},
"full_hidden": {0: "batch", 1: "seq_len"},
},
opset_version=17,
)
backbone_size = os.path.getsize(backbone_path)
print(f" Saved {backbone_path} ({backbone_size / 1e6:.1f} MB)")
# ─── Export lm_head: full_hidden → logits ─────────────────
print("Exporting lm_head (hidden → logits)...")
norm, lm_head = get_lm_head_and_norm(model)
if lm_head is None:
print(" WARNING: Could not find lm_head, skipping lm_head export")
return
head = GenericLMHead(norm, lm_head)
dummy_hidden = torch.randn(1, 32, D)
head_path = os.path.join(output_dir, "lm_head.onnx")
torch.onnx.export(
head, dummy_hidden,
head_path,
input_names=["full_hidden"],
output_names=["logits"],
dynamic_axes={
"full_hidden": {0: "batch", 1: "seq_len"},
"logits": {0: "batch", 1: "seq_len"},
},
opset_version=17,
)
head_size = os.path.getsize(head_path)
print(f" Saved {head_path} ({head_size / 1e6:.1f} MB)")
# Save tokenizer
tokenizer_path = os.path.join(output_dir, "tokenizer")
tokenizer.save_pretrained(tokenizer_path)
print(f" Saved {tokenizer_path}/")
# Verify
print("\nVerifying ONNX export...")
import onnxruntime as ort
sess_backbone = ort.InferenceSession(backbone_path)
sess_head = ort.InferenceSession(head_path)
test_text = "The capital of France is"
ids = tokenizer.encode(test_text, add_special_tokens=False)
input_ids = torch.tensor([ids]).numpy()
hidden_out, full_out = sess_backbone.run(None, {"input_ids": input_ids})
logits_out = sess_head.run(None, {"full_hidden": full_out})
next_token = logits_out[0][0, -1].argmax()
predicted = tokenizer.decode([next_token])
print(f" '{test_text}''{predicted}'")
print(f" Hidden shape: {hidden_out.shape}")
print(f" Logits shape: {logits_out[0].shape}")
print(f"\nDone. Total model size: {(backbone_size + head_size) / 1e6:.1f} MB")
print(f"Model: {model_name}")
if __name__ == "__main__":
parser = argparse.ArgumentParser(
description="Export a HuggingFace model as ONNX for epimem")
parser.add_argument("--model", type=str, default=DEFAULT_MODEL,
help=f"HuggingFace model name (default: {DEFAULT_MODEL})")
parser.add_argument("--output", type=str, default="models",
help="Output directory (default: models)")
args = parser.parse_args()
export(args.model, args.output)