Views
No views yet
prevalences/ - Datos de prevalencia de activaciones neuronales guardados durante el entrenamiento a intervalos regularesruns/ - Logs de TensorBoard y datos de ejecución del entrenamiento*.pth archivos - Checkpoints del modelo guardados en varios pasos de entrenamientosae_*.pth - Pesos finales del modelo SAE entrenado1/d_sae√(d_model/d_sae) ≈ 0.2041/d_model1 (renormalizadas cada paso)nn.Linear(2048, 49152, dtype=torch.float32) con biasnn.Linear(49152, 2048, dtype=torch.float32) con biasnn.Parameter de forma (49152,) inicializado con log(0.001)1def forward(self, x):
2 # x: (batch_size, 2048) - salidas de MLP de Llama
3 d = {}
4 original_input = x
5
6 # 1. Pre-procesamiento (centrado)
7 if self.use_pre_enc_bias:
8 x = x - self.dec.bias # (batch_size, 2048)
9
10 # 2. Encoding lineal
11 x = self.enc(x) # (batch_size, 49152)
12
13 # 3. Thresholding con función Step personalizada
14 threshold = torch.exp(self.log_threshold) # (49152,)
15 s = Step.apply(x, threshold) # (batch_size, 49152) - máscara binaria
16
17 # 4. Aplicar sparsity
18 x = x * s # (batch_size, 49152) - activaciones sparse
19
20 # 5. Decoding
21 x = self.dec(x) # (batch_size, 2048) - reconstrucción
22
23 # 6. Calcular métricas
24 d['mask'] = s # máscara de activaciones activas
25 d['reconstruction'] = ((x - original_input).pow(2)).mean(0).sum() # MSE loss
26
27 return d(x > threshold).to(x.dtype) - función escalón binariastep_*.pth - Checkpoints guardados cada 5,000 pasos después del paso 40,000sae_exp24_sparse0.001_d_sae_std_fullwarmup_steps256000_lr7e-05.pth - Modelo final entrenadoprevalences/step_*.npy - Datos de prevalencia de activaciones para análisis de neuronas muertas
prevalences/bin_edges.npy - Bordes de bins del histograma para análisis de prevalencia
pip install torch-tb-profiler tensorboard-plugin-profile1import torch
2from torch import nn
3from math import sqrt
4
5class Sae(nn.Module):
6 # ... implementación completa necesaria ...
7
8# Cargar modelo
9model = Sae(d_in=2048, d_sae=49152)
10model.load_state_dict(torch.load('sae_exp24_sparse0.001_d_sae_std_fullwarmup_steps256000_lr7e-05.pth'))
11model.eval()
12
13# Compilar para rendimiento óptimo
14model = torch.compile(model, mode="max-autotune")