gpt2-steering-denoiser / denoiser.py
electrostatickid's picture
Upload denoiser.py with huggingface_hub
c40763d verified
Raw
History Blame Contribute Delete
3.13 kB
"""
Denoiser для residual stream GPT-2 (слой 6, resid_post_mlp).
Conditioning на alpha (величину возмущения), а не абстрактный interpolation-t: обучение
использует ТУ ЖЕ форму искажения, что и инференс — аддитивную corrupted = h + alpha*u
(u — случайное единичное направление, НЕ настоящий steering-вектор v; сеть не должна знать
про валидационные v, но должна знать про типичный МАСШТАБ alpha). Conditioning-переменная —
alpha_norm = alpha / ALPHA_MAX_COND, подаётся явно и на train, и на inference с одинаковой
нормировкой — без этого denoiser physически не может научиться давать разную по силе
коррекцию для alpha=32 и alpha=256 (это и было причиной знакопеременных результатов раньше:
на инференсе передавался захардкоженный t=1 независимо от реального alpha).
corrupted = h + alpha*u, alpha ~ [0, ALPHA_MAX_COND], u ~ random unit vector
h_hat = denoiser(corrupted, alpha / ALPHA_MAX_COND)
L = || h - h_hat ||_2^2
На инференсе: h_tilde = denoiser(h + alpha*v, alpha / ALPHA_MAX_COND) — та же нормировка.
"""
import torch
import torch.nn as nn
ALPHA_MAX_COND = 300.0 # нормировка conditioning-переменной; должна совпадать train/inference
class Denoiser(nn.Module):
def __init__(self, d_model=768, hidden_mult=4):
super().__init__()
hidden = d_model * hidden_mult
self.net = nn.Sequential(
nn.Linear(d_model + 1, hidden),
nn.GELU(),
nn.Linear(hidden, hidden),
nn.GELU(),
nn.Linear(hidden, d_model),
)
def forward(self, h, alpha_norm=None):
# alpha_norm broadcast'ится под форму h с последней размерностью=1, независимо от
# ранга h (h может быть [batch, d_model] при обучении или [batch, pos, d_model] на
# инференсе во время генерации с KV-кешем).
target_shape = h.shape[:-1] + (1,)
if alpha_norm is None:
alpha_norm = torch.zeros(target_shape, device=h.device, dtype=h.dtype)
elif not torch.is_tensor(alpha_norm):
alpha_norm = torch.full(target_shape, float(alpha_norm), device=h.device, dtype=h.dtype)
else:
alpha_norm = alpha_norm.to(device=h.device, dtype=h.dtype).expand(target_shape)
# residual formulation: предсказываем поправку, а не сам h с нуля —
# это устойчивее при обучении и даёт identity-like поведение при alpha_norm~0
inp = torch.cat([h, alpha_norm], dim=-1)
return h + self.net(inp)