Spaces:
Paused
Paused
Download wave_tank_surgeon.py from Jibbalit/wave-tank-surgeon: direct link, hf CLI and curl.
- Browser
- Download file 34.6 kB
-
https://huggingface.co/spaces/Jibbalit/wave-tank-surgeon/resolve/main/wave_tank_surgeon.py
- Command line
-
hf download hf://spaces/Jibbalit/wave-tank-surgeon/wave_tank_surgeon.py
-
curl -L -o wave_tank_surgeon.py https://huggingface.co/spaces/Jibbalit/wave-tank-surgeon/resolve/main/wave_tank_surgeon.py
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() | |