Interpretabilidad mecanicista con TransformerLens

En un post pasado vimos una introducción a la interpretabilidad mecanicista y como intenta extraer información de los pesos y activaciones de los LLMs.

Una de las librerías más usadas en interpretabilidad mecanicista es TransformerLens, que permite cargar un LLM open source y explorar y transformar las activaciones del modelo.

Vamos a ver un ejemplo usando el modelo gpt2, siempre teniendo en cuenta que asumimos las hipótesis de la interpretabilidad mecanicista, especialmente que el modelo es un flujo residual lineal donde cada capa añade nueva información e interpretamos las activaciones de cada capa multiplicando por la matriz de embedding.

Después de instalar la librerías, cargamos el modelo, introducimos el texto de entrada y ejecutamos el “run_with_cache”, un método que realiza una pasada hacia adelante del modelo con la entrada definida y devuelve tanto los logits de salida como las activaciones.

import torch
from transformer_lens import HookedTransformer

# 1. Cargar el modelo desde Hugging Face usando TransformerLens
# 'gpt2' se descargará automáticamente de HF si no está en la caché local
print("Cargando el modelo...")
model = HookedTransformer.from_pretrained("gpt2")

model.cfg.use_attn_result = True

# 2. Definir un texto de prueba
prompt = "El gato se sentó en la"

# 3. Ejecutar el modelo y capturar las activaciones en el objeto 'cache'
# 'run_with_cache' devuelve las probabilidades de salida (logits) y la caché de activaciones
print("Ejecutando inferencia y capturando activaciones...")
logits, cache = model.run_with_cache(prompt)

# 4. Obtener la predicción del siguiente token
last_token_logits = logits[0, -1, :]
predicted_token_id = torch.argmax(last_token_logits).item()
predicted_token_str = model.to_string(predicted_token_id)

print(f"\nPrompt: '{prompt}'")
print(f"Siguiente token predicho: '{predicted_token_str}'")

# 5. Inspeccionar las activaciones internas guardadas en la caché
# Por ejemplo, las activaciones de la capa de atención (Attention Output) de la Capa 0
layer_0_attn_out = cache["blocks.0.attn.hook_result"]
print(f"\nForma del tensor de atención en la Capa 0: {layer_0_attn_out.shape}")
# Dimensiones: [lote, posición_del_token, número_de_cabezas, dimensión_de_la_cabeza]

Vemos como el siguiente token predicho es ” g” y el tensor de atención de la capa 0 tiene dimensiones (1, 9, 12, 768), ya que hay 9 tokens, 12 cabezales de atención y una dimensión de 768.

Cargando el modelo...
/tmp/ipykernel_1004/1663027743.py:7: DeprecationWarning: HookedTransformer.from_pretrained is deprecated and will be removed in a future major release. Use TransformerBridge.boot_transformers(...) instead, then call enable_compatibility_mode() for HookedTransformer-equivalent numerics. See docs/source/content/migrating_to_v3.md.
  model = HookedTransformer.from_pretrained("gpt2")
Loaded pretrained model gpt2 into HookedTransformer
Ejecutando inferencia y capturando activaciones...

Prompt: 'El gato se sentó en la'
Siguiente token predicho: ' g'

Forma del tensor de atención en la Capa 0: torch.Size([1, 9, 12, 768])

Ahora, obtenemos los pesos de atención del token “la” hacia todo el contexto en la capa 8 y el cabezal de atención 0.

capa = 8
cabezal = 0

# Obtener los pesos de atención de la última fila (el token " la") hacia todo el pasado
pesos_atencion = cache[f"blocks.{capa}.attn.hook_pattern"][0, cabezal, -1, :]

# 1. Añadimos [0] al final para obtener una lista plana de tokens
tokens = model.to_tokens(prompt)[0]

# 2. Ahora sí, decodificamos cada token individualmente
tokens_decodificados = [model.to_string(t) for t in tokens]

# 3. Tu bucle de impresión funcionará perfectamente ahora
for token, peso in zip(tokens_decodificados, pesos_atencion):
    print(f"Atención a {token:<12} : {peso.item():.4f}")

Observamos los pesos de atención y que este cabezal está “en reposo”. En la arquitectura Transformer, cuando un cabezal de atención calcula que la información semántica del prompt no coincide con su “especialidad”, prefiere no alterar el flujo residual. Como la función Softmax obliga a que todos los pesos de atención sumen 1.0, el cabezal deposita casi todo su peso en el primer token (el sumidero de atención) para actuar como un paso neutral.

Atención a <|endoftext|> : 0.8884
Atención a El           : 0.0041
Atención a  g           : 0.0056
Atención a ato          : 0.0111
Atención a  se          : 0.0058
Atención a  sent        : 0.0085
Atención a ó            : 0.0188
Atención a  en          : 0.0120
Atención a  la          : 0.0457

Ahora vamos a extraer el patrón de atención para el último token (“la”) de todos los cabezales de la capa 8 para ver si en todos el último token está centrándose en el primer token (sumidero) o si el cabezal está activo.

capa = 8
# Extraemos el patrón de atención de TODOS los cabezales de la capa 8 para la última posición
# Shape: [num_cabezales, key_posiciones]
patrones_capa = cache[f"blocks.{capa}.attn.hook_pattern"][0, :, -1, :]

print(f"--- ANÁLISIS DE LA CAPA {capa} ---")
for head_idx in range(model.cfg.n_heads):
    pesos_head = patrones_capa[head_idx]
    atencion_al_sumidero = pesos_head[0].item() # El índice 0 es <|endoftext|>

    # Si la atención al sumidero es menor al 50%, el cabezal está activo analizando el texto
    if atencion_al_sumidero < 0.50:
        # Encontremos a qué token del texto mira más (excluyendo el sumidero)
        token_max_idx = pesos_head[1:].argmax().item() + 1
        token_max_texto = tokens_decodificados[token_max_idx]
        peso_max = pesos_head[token_max_idx].item()

        print(f"¡Cabezal {capa}.{head_idx} ACTIVO! Atiende principalmente a '{token_max_texto}' con {peso_max:.4f}")
    else:
        print(f"Cabezal {capa}.{head_idx}: En reposo (Sumidero: {atencion_al_sumidero:.4f})")
     

Observamos que se han encontrado dos cabezales activos, el 8.6 que se atiende a sí mismo y el 8.11 que atiende a “ato”. Se ha encontrado un circuito de procesamiento lingüístico. En las capas siguientes (9, 10 y 11), el modelo va a fusionar la información de estos dos cabezales mediante una operación matemática interna: Restricción de 8.6 (femenino) + información de 8.11 (gato).

--- ANÁLISIS DE LA CAPA 8 ---
Cabezal 8.0: En reposo (Sumidero: 0.8884)
Cabezal 8.1: En reposo (Sumidero: 0.7933)
Cabezal 8.2: En reposo (Sumidero: 0.8977)
Cabezal 8.3: En reposo (Sumidero: 0.9089)
Cabezal 8.4: En reposo (Sumidero: 0.9839)
Cabezal 8.5: En reposo (Sumidero: 0.6162)
¡Cabezal 8.6 ACTIVO! Atiende principalmente a ' la' con 0.1831
Cabezal 8.7: En reposo (Sumidero: 0.5314)
Cabezal 8.8: En reposo (Sumidero: 0.8848)
Cabezal 8.9: En reposo (Sumidero: 0.9301)
Cabezal 8.10: En reposo (Sumidero: 0.6700)
¡Cabezal 8.11 ACTIVO! Atiende principalmente a 'ato' con 0.1652

Ahora vamos a comprobar si la salida del cabezal 11 de la capa 8 en el último token tiene impacto en ” g” y por lo tanto el siguiente token tiene en cuenta la información “ato” (gato). Para ello, multiplicamos el vector de salida del cabezal 8.11 por la matriz de desembedding W_U.

# 1. Asegurar la configuración
model.cfg.use_attn_result = True
logits, cache = model.run_with_cache(prompt)

# 2. Obtener el ID de tu token objetivo (' g')
target_token_id = model.to_single_token(" g")

# 3. Extraer el resultado de ESCRITURA de la capa 8
# Shape: [batch, posicion, cabezal, d_model]
layer_8_results = cache["blocks.8.attn.hook_result"]

# 4. Aislar el impacto del Cabezal 11 en el último token (posición -1) hacia ' g'
# Multiplicamos el vector de salida del cabezal 8.11 por la matriz de desempaquetado W_U
head_8_11_vector = layer_8_results[0, -1, 11, :]
contribucion_logit = head_8_11_vector @ model.W_U[:, target_token_id]

print(f"--- ATRIBUCIÓN DIRECTA DEL CABEZAL 8.11 ---")
print(f"Impacto directo en el logit de ' g': {contribucion_logit.item():.4f}")

Como sale un resultado positivo, el cabezal 8.11 sí tiene atribución directa.

--- ATRIBUCIÓN DIRECTA DEL CABEZAL 8.11 ---
Impacto directo en el logit de ' g': 0.7364

Hemos visto como la librería TranformerLens puede ser muy útil para tener una interpretabilidad mecanicista de los LLMs siempre que asumamos sus hipótesis.

Deja una respuesta

Orgullosamente ofrecido por WordPress | Tema: Baskerville 2 por Anders Noren.

Subir ↑