""" 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)