wave-tank-surgeon / wave_tank_surgeon.py
Jibbalit's picture
Fix attention IO capture for keyword-only Llama modules
1f194e1
Raw History Blame Contribute Delete
34.6 kB
#!/usr/bin/env python3
# /// script
# dependencies = [
# "torch",
# "transformers",
# "accelerate",
# "huggingface_hub",
# ]
# ///
"""
wave_tank_surgeon.py - Replace FFN and Attention layers in ANY HuggingFace model
with wave tank processors. Scans model architecture, trains wave tank replacements
to mimic each layer's behavior, swaps them in, saves the surgically modified model.
Usage:
python wave_tank_surgeon.py --model deepseek-ai/DeepSeek-R1-Distill-Qwen-14B
python wave_tank_surgeon.py --model ./my_local_model --ffn-grid 16 --attn-grid 4
python wave_tank_surgeon.py --model gpt2 --steps 500 --device cuda:0
Architecture (verified on DeepSeek-R1-Distill-Qwen-14B, 94.3% size reduction):
- WaveTankFFN: grid=16, 98.8% param reduction, converges step ~200
- WaveTankAttention: grid=4 per-head, 79.2% param reduction, converges step ~300
- StatelessWaveTank core: 2D wave equation via 5-point stencil Laplacian
"""
import argparse
import json
import os
import sys
import time
import math
from collections import OrderedDict
from typing import Dict, List, Optional, Tuple
import torch
import torch.nn as nn
import torch.nn.functional as F
from torch.utils.data import DataLoader, TensorDataset
# ============================================================================
# WAVE TANK CORE
# ============================================================================
class StatelessWaveTank(nn.Module):
"""2D wave equation processor. Nonlinear compute primitive.
∂²u/∂t² = c² ∇²u + S(x,y,t)
Uses 5-point stencil Laplacian with learned alpha (wave speed) and damping.
Runs n_steps of wave dynamics on a grid_size x grid_size excitation field.
"""
def __init__(self, grid_size=16, wave_alpha=0.1, n_steps=3, damping=0.99):
super().__init__()
self.grid_size = grid_size
self.n_steps = n_steps
self.alpha = nn.Parameter(torch.tensor(wave_alpha))
self.damping = nn.Parameter(torch.tensor(damping))
def forward(self, excitation):
# excitation: [..., grid_size, grid_size]
u = torch.zeros_like(excitation)
u_prev = torch.zeros_like(excitation)
for _ in range(self.n_steps):
laplacian = (
torch.roll(u_prev, 1, -2) + torch.roll(u_prev, -1, -2) +
torch.roll(u_prev, 1, -1) + torch.roll(u_prev, -1, -1) -
4 * u_prev
)
u_next = self.damping * (2 * u - u_prev + self.alpha * laplacian) + excitation
u_prev = u
u = u_next
return u
# ============================================================================
# WAVE TANK FFN REPLACEMENT
# ============================================================================
class WaveTankFFN(nn.Module):
"""Drop-in FFN replacement using wave tank processor.
input → [Linear proj d_model → grid*grid] → [WaveTank] → [Linear map grid*grid → d_model] → output
Grid=16 (256 cells) is REQUIRED for FFN — grid=4 fails to approximate
high-dimensional SwiGLU hidden states.
"""
def __init__(self, d_model, d_ff=None, grid_size=16, n_steps=3):
super().__init__()
self.d_model = d_model
self.grid_size = grid_size
self.grid_dim = grid_size * grid_size
self.in_proj = nn.Linear(d_model, self.grid_dim, bias=False)
self.wave_tank = StatelessWaveTank(grid_size=grid_size, n_steps=n_steps)
self.out_proj = nn.Linear(self.grid_dim, d_model, bias=False)
self.act = nn.SiLU()
def forward(self, x):
B, S, D = x.shape
N = B * S
h = self.act(self.in_proj(x)) # [B*S, grid_dim]
h = h.view(N, self.grid_size, self.grid_size) # [N, G, G]
h = self.wave_tank(h) # [N, G, G]
h = h.view(B, S, self.grid_dim) # [B, S, grid_dim]
return self.out_proj(h) # [B, S, d_model]
# ============================================================================
# WAVE TANK ATTENTION REPLACEMENT
# ============================================================================
class WaveTankAttention(nn.Module):
"""Drop-in attention replacement using wave tank processor.
Per-head Q/K/V projections into grid space are MANDATORY.
Shared grid across heads fails (loss plateaus at 0.749).
Grid=4 per-head works (79.2% reduction). Grid=8 also works for more capacity.
Flow:
x → per-head Q/K/V proj → reshape to [N, H, G, G] → wave tank each
→ compute attention scores via einsum → softmax → weighted sum → out_proj
"""
def __init__(self, d_model, n_heads, grid_size=4, n_steps=3):
super().__init__()
self.d_model = d_model
self.n_heads = n_heads
self.grid_size = grid_size
self.grid_dim = grid_size * grid_size
self.head_dim = self.grid_dim # each head maps to grid
self.q_proj = nn.Linear(d_model, n_heads * self.grid_dim, bias=False)
self.k_proj = nn.Linear(d_model, n_heads * self.grid_dim, bias=False)
self.v_proj = nn.Linear(d_model, n_heads * self.grid_dim, bias=False)
self.wave_tank = StatelessWaveTank(grid_size=grid_size, n_steps=n_steps)
self.out_proj = nn.Linear(n_heads * self.grid_dim, d_model, bias=False)
self.act = nn.SiLU()
self.norm = nn.LayerNorm(n_heads * self.grid_dim)
# Causal mask will be created dynamically
self._causal_mask = None
self._causal_mask_size = 0
self.return_tuple = False
self.output_tuple_len = 0
def _get_causal_mask(self, S, device, dtype):
if self._causal_mask is None or self._causal_mask_size < S:
self._causal_mask_size = S
mask = torch.triu(torch.full((S, S), float('-inf'), dtype=dtype, device=device), diagonal=1)
self._causal_mask = mask
return self._causal_mask[:S, :S]
def forward(self, x=None, attention_mask=None, *args, **kwargs):
if x is None:
x = kwargs.get('hidden_states')
if x is None and args:
x = args[0]
if x is None:
raise ValueError("WaveTankAttention requires hidden_states as positional x or keyword hidden_states")
if attention_mask is None:
attention_mask = kwargs.get('attention_mask')
B, S, D = x.shape
gs = self.grid_size
H = self.n_heads
gd = self.grid_dim
N = B * S
Q = self.act(self.q_proj(x)).view(N, H, gs, gs)
K = self.act(self.k_proj(x)).view(N, H, gs, gs)
V = self.act(self.v_proj(x)).view(N, H, gs, gs)
# Wave tank processes each head's Q, K, V
Q_w = self.wave_tank(Q).view(B, S, H, gd)
K_w = self.wave_tank(K).view(B, S, H, gd)
V_w = self.wave_tank(V).view(B, S, H, gd)
# Attention scores: [B, S, T, H] where T=S (self-attention)
scores = torch.einsum('bshd,bthd->bsth', Q_w, K_w) / math.sqrt(gd)
# Causal mask
causal = self._get_causal_mask(S, x.device, x.dtype)
scores = scores + causal.unsqueeze(0).unsqueeze(-1) # [1, S, S, 1]
# Softmax over source sequence dimension (dim=2)
weights = F.softmax(scores, dim=2)
# Weighted sum
out = torch.einsum('bsth,bthd->bshd', weights, V_w) # [B, S, H, gd]
out = self.norm(self.act(out.reshape(B, S, H * gd)))
out = self.out_proj(out)
if self.return_tuple:
tuple_len = max(self.output_tuple_len, 1)
return tuple([out] + [None] * (tuple_len - 1))
return out
# ============================================================================
# MODEL SCANNER - Detect FFN and Attention layers in any HF model
# ============================================================================
class LayerInfo:
def __init__(self, path, layer_type, d_model, d_ff=None, n_heads=None, layer=None):
self.path = path
self.layer_type = layer_type # 'ffn' or 'attention'
self.d_model = d_model
self.d_ff = d_ff
self.n_heads = n_heads
self.layer = layer # reference to the actual nn.Module
self.output_is_tuple = False
self.output_tuple_len = 0
def __repr__(self):
extra = f", d_ff={self.d_ff}" if self.d_ff else ""
extra += f", n_heads={self.n_heads}" if self.n_heads else ""
return f"LayerInfo({self.path}, {self.layer_type}, d_model={self.d_model}{extra})"
def scan_model(model) -> List[LayerInfo]:
"""Scan any HuggingFace model and identify FFN and attention layers.
Works with:
- GPT-2 family (Conv1D based)
- LLaMA / Mistral / Qwen family (MLP + Attention)
- DeepSeek family (MLP + Attention)
- Gemma family
- Most other transformer architectures
Returns list of LayerInfo objects with enough metadata to create replacements.
"""
layers = []
name_to_module = dict(model.named_modules())
# Strategy: walk named modules, detect known patterns
# Track which modules are parents of detected layers to avoid double-counting
detected_paths = set()
all_modules = list(model.named_modules())
# Sort by path length (most specific first) so we detect leaf layers before parents
all_modules.sort(key=lambda x: x[0].count('.'), reverse=True)
for name, module in all_modules:
if not name: # skip root
continue
# Skip if this module is a parent of an already-detected layer
if any(p.startswith(name + '.') for p in detected_paths):
continue
cls_name = type(module).__name__.lower()
# --- Detect FFN / MLP layers ---
if _is_ffn(module, cls_name):
info = _extract_ffn_info(name, module, model)
if info:
layers.append(info)
detected_paths.add(name)
continue
# --- Detect Attention layers ---
if _is_attention(module, cls_name):
info = _extract_attn_info(name, module, model)
if info:
layers.append(info)
detected_paths.add(name)
continue
return layers
def _is_ffn(module, cls_name):
"""Check if module is an FFN/MLP block."""
ffn_names = {'mlp', 'ffn', 'feedforward', 'feed_forward'}
if cls_name in ffn_names:
return True
# Check for typical FFN structure: has 2+ linear layers and activation
if hasattr(module, 'gate_proj') or hasattr(module, 'up_proj') or hasattr(module, 'down_proj'):
return True
if hasattr(module, 'fc1') and hasattr(module, 'fc2'):
return True
if hasattr(module, 'c_fc') and hasattr(module, 'c_proj'):
return True # GPT-2 style
return False
def _is_attention(module, cls_name):
"""Check if module is an attention block."""
attn_names = {'attention', 'selfattention', 'self_attention', 'sdpaattention',
'eagerattention', 'gqaattention', 'attentionblock'}
if cls_name in attn_names:
return True
# Reject parent decoder layers that merely contain attention
if hasattr(module, 'self_attn') and (hasattr(module, 'mlp') or hasattr(module, 'ffn')):
return False
if hasattr(module, 'attn') and (hasattr(module, 'mlp') or hasattr(module, 'ffn')):
return False
# Has Q/K/V or Q+K+V combined projection
if (hasattr(module, 'q_proj') or hasattr(module, 'query') or
hasattr(module, 'c_attn')):
return True
# GPT-2 Conv1D attention
if hasattr(module, 'c_attn') and hasattr(module, 'c_proj'):
return True
return False
def _extract_ffn_info(name, module, model):
"""Extract FFN metadata from the module."""
d_model = None
d_ff = None
# LLaMA/Mistral/DeepSeek style: gate_proj, up_proj, down_proj (SwiGLU)
if hasattr(module, 'gate_proj'):
d_model = module.down_proj.out_features if hasattr(module.down_proj, 'out_features') else None
d_ff = module.gate_proj.out_features if hasattr(module.gate_proj, 'out_features') else None
if d_model is None:
# Try to infer from weight shape
w = module.down_proj.weight
d_model = w.shape[0] if w.shape[0] < w.shape[1] else w.shape[1]
return LayerInfo(name, 'ffn', d_model, d_ff=d_ff, layer=module)
# GPT-2 style: c_fc, c_proj (Conv1D)
if hasattr(module, 'c_fc') and hasattr(module, 'c_proj'):
# Conv1D weight shape: [in_features, out_features]
# c_fc: [d_model, d_ff], c_proj: [d_ff, d_model]
d_model = module.c_fc.weight.shape[0]
d_ff = module.c_fc.weight.shape[1]
return LayerInfo(name, 'ffn', d_model, d_ff=d_ff, layer=module)
# Standard style: fc1, fc2 or similar
if hasattr(module, 'fc1') and hasattr(module, 'fc2'):
d_model = module.fc2.out_features if hasattr(module.fc2, 'out_features') else None
d_ff = module.fc1.out_features if hasattr(module.fc1, 'out_features') else None
return LayerInfo(name, 'ffn', d_model, d_ff=d_ff, layer=module)
# Generic: look for linear layers inside
linears = [(n, m) for n, m in module.named_modules() if isinstance(m, nn.Linear)]
if len(linears) >= 2:
# First linear: d_model → d_ff, last linear: d_ff → d_model
first = linears[0][1]
last = linears[-1][1]
d_ff = first.out_features
d_model = last.out_features
return LayerInfo(name, 'ffn', d_model, d_ff=d_ff, layer=module)
return None
def _extract_attn_info(name, module, model):
"""Extract attention metadata from the module."""
d_model = None
n_heads = None
# Try common attribute names for number of heads
for attr in ['num_heads', 'n_heads', 'num_attention_heads']:
if hasattr(module, attr):
n_heads = getattr(module, attr)
break
# Try to get d_model from config or projection dimensions
for attr in ['embed_dim', 'hidden_size', 'd_model']:
if hasattr(module, attr):
d_model = getattr(module, attr)
break
# Infer from weight shapes
if d_model is None:
if hasattr(module, 'q_proj') and isinstance(module.q_proj, nn.Linear):
d_model = module.q_proj.in_features
elif hasattr(module, 'c_attn'): # GPT-2 Conv1D
d_model = module.c_attn.weight.shape[1]
elif hasattr(module, 'query') and isinstance(module.query, nn.Linear):
d_model = module.query.in_features
# Infer n_heads if not found
if n_heads is None:
if d_model is not None:
# Default: assume head_dim=64 or 128
if d_model % 128 == 0:
n_heads = d_model // 128
elif d_model % 64 == 0:
n_heads = d_model // 64
else:
n_heads = max(1, d_model // 96)
if d_model is not None and n_heads is not None:
return LayerInfo(name, 'attention', d_model, n_heads=n_heads, layer=module)
return None
# ============================================================================
# TRAINING: Extract IO pairs from original layers, train wave tank replacements
# ============================================================================
def collect_io_pairs(layer_info, model, dataloader, device, max_batches=50):
"""Run forward passes through the model, hooking into the target layer
to collect input/output pairs for training the wave tank replacement."""
inputs_list = []
outputs_list = []
layer = layer_info.layer
path = layer_info.path
hook_registered = False
captured = {}
def first_tensor(value):
if torch.is_tensor(value):
return value
if isinstance(value, dict):
for key in ('hidden_states', 'last_hidden_state'):
tensor = value.get(key)
if torch.is_tensor(tensor):
return tensor
for item in value.values():
tensor = first_tensor(item)
if tensor is not None:
return tensor
if isinstance(value, (tuple, list)):
for item in value:
tensor = first_tensor(item)
if tensor is not None:
return tensor
if hasattr(value, 'to_tuple'):
return first_tensor(value.to_tuple())
return None
def hook_fn(module, input_args, input_kwargs, output):
input_tensor = first_tensor(input_args)
if input_tensor is None:
input_tensor = first_tensor(input_kwargs)
output_tensor = first_tensor(output)
if input_tensor is None or output_tensor is None:
return
captured['input'] = input_tensor.detach()
captured['output'] = output_tensor.detach()
captured['output_is_tuple'] = isinstance(output, tuple)
captured['output_tuple_len'] = len(output) if isinstance(output, tuple) else 0
try:
handle = layer.register_forward_hook(hook_fn, with_kwargs=True)
except TypeError:
def legacy_hook_fn(module, input_args, output):
hook_fn(module, input_args, {}, output)
handle = layer.register_forward_hook(legacy_hook_fn)
hook_registered = True
model.eval()
with torch.no_grad():
for batch_idx, batch in enumerate(dataloader):
if batch_idx >= max_batches:
break
try:
# Move batch to device
if isinstance(batch, dict):
batch = {k: v.to(device) for k, v in batch.items()}
model(**batch)
elif isinstance(batch, (list, tuple)):
batch = tuple(v.to(device) for v in batch)
model(*batch)
else:
batch = batch.to(device)
model(batch)
if 'input' in captured and 'output' in captured:
layer_info.output_is_tuple = captured.get('output_is_tuple', False)
layer_info.output_tuple_len = captured.get('output_tuple_len', 0)
inputs_list.append(captured['input'].cpu())
outputs_list.append(captured['output'].cpu())
captured.clear()
except Exception as e:
print(f" Warning: batch {batch_idx} failed: {e}")
continue
if hook_registered:
handle.remove()
if not inputs_list:
print(f" ERROR: No IO pairs collected for {path}")
return None, None
all_inputs = torch.cat(inputs_list, dim=0)
all_outputs = torch.cat(outputs_list, dim=0)
print(f" Collected {all_inputs.shape[0]} samples, input shape {all_inputs.shape}, output shape {all_outputs.shape}")
return all_inputs, all_outputs
def train_replacement(layer_info, inputs, outputs, device, ffn_grid=16, attn_grid=4,
steps=5000, lr=1e-3, batch_size=32):
"""Train a wave tank replacement to mimic the original layer's behavior."""
path = layer_info.path
d_model = layer_info.d_model
layer_type = layer_info.layer_type
target_dtype = inputs.dtype if inputs.is_floating_point() else torch.float32
train_dtype = torch.float32
# Create replacement module
if layer_type == 'ffn':
replacement = WaveTankFFN(d_model=d_model, grid_size=ffn_grid, n_steps=3).to(
device=device,
dtype=train_dtype,
)
label = f"WaveTankFFN(grid={ffn_grid})"
elif layer_type == 'attention':
n_heads = layer_info.n_heads
replacement = WaveTankAttention(d_model=d_model, n_heads=n_heads,
grid_size=attn_grid, n_steps=3).to(
device=device,
dtype=train_dtype,
)
label = f"WaveTankAttn(grid={attn_grid}, heads={n_heads})"
else:
return None
# Count params
orig_params = sum(p.numel() for p in layer_info.layer.parameters())
repl_params = sum(p.numel() for p in replacement.parameters())
reduction = 1.0 - repl_params / max(orig_params, 1)
print(f" {label}: {repl_params:,} params vs {orig_params:,} original ({reduction*100:.1f}% reduction)")
# Prepare data
dataset = TensorDataset(inputs, outputs)
loader = DataLoader(dataset, batch_size=batch_size, shuffle=True, drop_last=True)
optimizer = torch.optim.Adam(replacement.parameters(), lr=lr)
loss_fn = nn.MSELoss()
# Flatten outputs for FFN (attention outputs may have different shapes)
# Ensure inputs/outputs are [batch, seq, d_model]
if inputs.dim() == 2:
inputs = inputs.unsqueeze(1)
outputs = outputs.unsqueeze(1)
replacement.train()
best_loss = float('inf')
best_state = None
for step in range(steps):
epoch_loss = 0.0
n_batches = 0
for batch_in, batch_out in loader:
batch_in = batch_in.to(device=device, dtype=train_dtype)
batch_out = batch_out.to(device=device, dtype=train_dtype)
optimizer.zero_grad()
if layer_type == 'ffn':
pred = replacement(batch_in)
else:
# Attention: may need to handle residual vs direct output
try:
pred = replacement(batch_in)
except Exception:
continue
loss = loss_fn(pred, batch_out)
loss.backward()
# Gradient clipping for stability
torch.nn.utils.clip_grad_norm_(replacement.parameters(), 1.0)
optimizer.step()
epoch_loss += loss.item()
n_batches += 1
avg_loss = epoch_loss / max(n_batches, 1)
if avg_loss < best_loss:
best_loss = avg_loss
best_state = {k: v.clone() for k, v in replacement.state_dict().items()}
if (step + 1) % 50 == 0 or step == 0:
alpha = replacement.wave_tank.alpha.item()
damping = replacement.wave_tank.damping.item()
print(f" Step {step+1}/{steps}: loss={avg_loss:.6f} (best={best_loss:.6f}) "
f"alpha={alpha:.4f} damping={damping:.4f}")
# Load best weights
if best_state is not None:
replacement.load_state_dict(best_state)
if layer_type == 'attention':
replacement.return_tuple = layer_info.output_is_tuple
replacement.output_tuple_len = layer_info.output_tuple_len
return replacement.to(device=device, dtype=target_dtype)
# ============================================================================
# SURGERY: Replace layers in the model
# ============================================================================
def perform_surgery(model, layer_info, replacement):
"""Replace a layer in the model with its wave tank replacement.
Navigates the module hierarchy using dot-separated path and swaps the module.
"""
path = layer_info.path
parts = path.split('.')
# Navigate to parent module
parent = model
for part in parts[:-1]:
if hasattr(parent, part):
parent = getattr(parent, part)
elif isinstance(parent, nn.ModuleList) or isinstance(parent, nn.ModuleDict):
parent = parent[int(part)] if part.isdigit() else parent[part]
else:
raise ValueError(f"Cannot navigate to {path}: stuck at {'.'.join(parts[:parts.index(part)])}")
final_name = parts[-1]
# Set the replacement
if isinstance(parent, nn.ModuleList):
setattr(parent, final_name, replacement) # ModuleList items set via __setattr__
elif isinstance(parent, nn.ModuleDict):
parent[final_name] = replacement
else:
setattr(parent, final_name, replacement)
print(f" Surgery: replaced {path} with {type(replacement).__name__}")
# ============================================================================
# MAIN: Full pipeline
# ============================================================================
def main():
parser = argparse.ArgumentParser(description="Wave Tank Surgeon - Replace FFN/Attention in any model")
parser.add_argument('--model', required=True, help='HuggingFace model name or local path')
parser.add_argument('--output', default=None, help='Output directory for replaced model')
parser.add_argument('--ffn-grid', type=int, default=16, help='Grid size for FFN replacement (default: 16)')
parser.add_argument('--attn-grid', type=int, default=16, help='Grid size for attention replacement (default: 4)')
parser.add_argument('--steps', type=int, default=5000, help='Training steps per layer (default: 500)')
parser.add_argument('--lr', type=float, default=1e-3, help='Learning rate (default: 1e-3)')
parser.add_argument('--batch-size', type=int, default=16, help='Batch size for IO collection (default: 4)')
parser.add_argument('--max-batches', type=int, default=500, help='Max batches for IO pair collection (default: 50)')
parser.add_argument('--device', default=None, help='Device (default: auto-detect)')
parser.add_argument('--skip-attn', action='store_true', help='Skip attention replacement (FFN only)')
parser.add_argument('--skip-ffn', action='store_true', help='Skip FFN replacement (attention only)')
parser.add_argument('--dtype', default='float16', help='Model dtype: float16, float32, bfloat16')
parser.add_argument('--dry-run', action='store_true', help='Scan model and report layers, but do not train')
parser.add_argument('--push-to-hub', default=None,
help='Push replaced model to HF Hub (e.g. username/model-wave-tank)')
parser.add_argument('--private', action='store_true', help='Make pushed Hub repo private')
args = parser.parse_args()
# Device
if args.device:
device = torch.device(args.device)
elif torch.cuda.is_available():
device = torch.device('cuda:0')
else:
device = torch.device('cpu')
# Dtype
dtype_map = {
'float16': torch.float16, 'fp16': torch.float16,
'float32': torch.float32, 'fp32': torch.float32,
'bfloat16': torch.bfloat16, 'bf16': torch.bfloat16,
}
torch_dtype = dtype_map.get(args.dtype, torch.float16)
print("=" * 70)
print("WAVE TANK SURGEON")
print(f"Model: {args.model}")
print(f"Device: {device} | Dtype: {torch_dtype}")
print(f"FFN grid: {args.ffn_grid} | Attn grid: {args.attn_grid}")
print(f"Training: {args.steps} steps, lr={args.lr}")
print("=" * 70)
hf_token = (
os.environ.get('HF_TOKEN')
or os.environ.get('HUGGING_FACE_HUB_TOKEN')
or os.environ.get('HUGGINGFACEHUB_API_TOKEN')
)
# Load model and tokenizer
print(f"Loading model: {args.model}...")
from transformers import AutoModelForCausalLM, AutoTokenizer
tokenizer = AutoTokenizer.from_pretrained(
args.model,
trust_remote_code=True,
token=hf_token,
)
model = AutoModelForCausalLM.from_pretrained(
args.model,
torch_dtype=torch_dtype,
device_map='auto' if torch.cuda.is_available() else None,
trust_remote_code=True,
token=hf_token,
)
if device.type == 'cpu' or (device.type == 'cuda' and not hasattr(model, 'hf_device_map')):
model = model.to(device)
print(f"Model loaded. Total params: {sum(p.numel() for p in model.parameters()):,}")
# Scan model
print(f"Scanning model architecture...")
layer_infos = scan_model(model)
ffn_layers = [l for l in layer_infos if l.layer_type == 'ffn']
attn_layers = [l for l in layer_infos if l.layer_type == 'attention']
print(f"Found {len(ffn_layers)} FFN layers and {len(attn_layers)} attention layers")
for l in layer_infos:
print(f" {l}")
if args.dry_run:
print("Dry run: exiting after scan.")
return model
# Prepare calibration data (use tokenizer to create random inputs)
print(f"Preparing calibration data...")
calib_texts = [
"The wave tank computes interference patterns through the discrete Laplacian.",
"Attention mechanisms allow models to focus on relevant parts of the input.",
"The transformer architecture has revolutionized natural language processing.",
"Wave equations describe the propagation of disturbances through a medium.",
"Deep learning models can be compressed using knowledge distillation.",
"The Laplacian operator measures the divergence of the gradient field.",
"Neural network pruning removes redundant parameters while preserving accuracy.",
"Quantization reduces the precision of weights to decrease model size.",
]
encoded = tokenizer(
calib_texts,
return_tensors='pt',
padding=True,
truncation=True,
max_length=128,
)
# Create dataloader from calibration data
input_ids = encoded['input_ids']
attention_mask = encoded['attention_mask']
# Repeat to get enough batches
reps = max(args.max_batches, 1)
input_ids = input_ids.repeat(reps, 1)
attention_mask = attention_mask.repeat(reps, 1)
batch_size = min(args.batch_size, input_ids.shape[0])
calib_loader = DataLoader(
TensorDataset(input_ids, attention_mask),
batch_size=batch_size,
shuffle=False,
)
# Convert to dict-style loader for HuggingFace models
class DictLoader:
def __init__(self, input_ids, attention_mask, batch_size):
self.input_ids = input_ids
self.attention_mask = attention_mask
self.batch_size = batch_size
self.n = input_ids.shape[0]
def __iter__(self):
for i in range(0, self.n, self.batch_size):
yield {
'input_ids': self.input_ids[i:i+self.batch_size],
'attention_mask': self.attention_mask[i:i+self.batch_size],
}
def __len__(self):
return (self.n + self.batch_size - 1) // self.batch_size
dict_loader = DictLoader(input_ids, attention_mask, batch_size)
# Train and replace each layer
replacements = {}
layers_to_process = []
if not args.skip_ffn:
layers_to_process.extend(ffn_layers)
if not args.skip_attn:
layers_to_process.extend(attn_layers)
# Sort by path for deterministic order
layers_to_process.sort(key=lambda l: l.path)
for i, layer_info in enumerate(layers_to_process):
print(f"")
print(f"[{i+1}/{len(layers_to_process)}] Processing {layer_info.path} ({layer_info.layer_type})")
print(f" d_model={layer_info.d_model}" +
(f", d_ff={layer_info.d_ff}" if layer_info.d_ff else "") +
(f", n_heads={layer_info.n_heads}" if layer_info.n_heads else ""))
# Collect IO pairs
print(f" Collecting IO pairs...")
inputs, outputs = collect_io_pairs(
layer_info, model, dict_loader, device,
max_batches=args.max_batches
)
if inputs is None:
print(f" SKIPPING: could not collect IO pairs")
continue
# Reshape if needed - ensure [batch, seq, d_model]
if inputs.dim() == 2:
inputs = inputs.unsqueeze(1)
if outputs.dim() == 2:
outputs = outputs.unsqueeze(1)
# For attention layers, outputs might be tuple - just use first element
# (the actual attention output, not past_key_values etc.)
# Train replacement
grid = args.ffn_grid if layer_info.layer_type == 'ffn' else args.attn_grid
replacement = train_replacement(
layer_info, inputs, outputs, device,
ffn_grid=args.ffn_grid, attn_grid=args.attn_grid,
steps=args.steps, lr=args.lr,
)
if replacement is not None:
replacements[layer_info.path] = replacement
# Replace in model
replacement.eval()
perform_surgery(model, layer_info, replacement)
# Report
print(f"")
print("=" * 70)
print("SURGERY COMPLETE")
orig_size = 0
new_size = 0
for l in layer_infos:
op = sum(p.numel() for p in l.layer.parameters())
orig_size += op
for p in model.parameters():
new_size += p.numel()
print(f"Original total: {orig_size:,} params in replaced layers")
print(f"Model now: {sum(p.numel() for p in model.parameters()):,} total params")
print(f"Replaced {len(replacements)}/{len(layers_to_process)} layers")
print("=" * 70)
# Save
if args.output:
output_dir = args.output
else:
model_name = args.model.replace('/', '_')
output_dir = f"{model_name}_wave_tank"
os.makedirs(output_dir, exist_ok=True)
print(f"Saving replaced model to {output_dir}...")
model.save_pretrained(output_dir)
tokenizer.save_pretrained(output_dir)
# Save surgery report
report = {
'source_model': args.model,
'ffn_grid': args.ffn_grid,
'attn_grid': args.attn_grid,
'training_steps': args.steps,
'replaced_layers': {p: type(r).__name__ for p, r in replacements.items()},
'skipped_layers': [l.path for l in layers_to_process if l.path not in replacements],
'ffn_count': len([p for p in replacements.values() if isinstance(p, WaveTankFFN)]),
'attn_count': len([p for p in replacements.values() if isinstance(p, WaveTankAttention)]),
}
with open(os.path.join(output_dir, 'surgery_report.json'), 'w') as f:
json.dump(report, f, indent=2)
print(f"Saved. Surgery report: {os.path.join(output_dir, 'surgery_report.json')}")
# Push to Hub if requested
if args.push_to_hub:
print(f"Pushing to HuggingFace Hub: {args.push_to_hub}...")
model.push_to_hub(args.push_to_hub, private=args.private, token=hf_token)
tokenizer.push_to_hub(args.push_to_hub, private=args.private, token=hf_token)
# Also push surgery report
from huggingface_hub import HfApi
api = HfApi(token=hf_token)
api.upload_file(
path_or_fileobj=os.path.join(output_dir, 'surgery_report.json'),
path_in_repo='surgery_report.json',
repo_id=args.push_to_hub,
repo_type='model',
)
print(f"Pushed to Hub: https://huggingface.co/{args.push_to_hub}")
return model
if __name__ == '__main__':
main()