| """ |
| 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 |
|
|
|
|
| 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): |
| |
| |
| |
| 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) |
| |
| |
| inp = torch.cat([h, alpha_norm], dim=-1) |
| return h + self.net(inp) |
|
|