¿Qué es la máscara de relleno de clave frente a la máscara de atención en PyTorch?
sobes.tech AI
Respuesta de la IA
En PyTorch, al trabajar con transformadores y atención, a menudo se utilizan máscaras para el procesamiento correcto de secuencias de diferentes longitudes y para evitar la atención a posiciones no deseadas.
-
Máscara de relleno clave (Key Padding Mask): es una máscara que indica qué posiciones en la secuencia son relleno (rellenos), para que el modelo no las tenga en cuenta al calcular la atención. Normalmente, es un tensor booleano donde True significa posición de relleno. Se aplica a las claves (keys) y valores (values) en el mecanismo de atención para ignorar los tokens de relleno.
-
Máscara de atención (Attention Mask): una máscara más general que controla qué posiciones pueden influir en la posición actual al calcular la atención. Por ejemplo, en tareas de generación automática de texto, se usa una máscara que prohíbe mirar hacia adelante (tokens futuros) para evitar la fuga de información. Esta máscara puede ser una matriz triangular, donde ciertas conexiones están prohibidas.
Ejemplo:
import torch
from torch.nn import MultiheadAttention
# Supongamos que tenemos una secuencia con relleno
key_padding_mask = torch.tensor([[False, False, True, True]]) # True para relleno
# Máscara de atención para evitar atención a tokens futuros
attn_mask = torch.triu(torch.ones(4, 4), diagonal=1).bool() # Máscara triangular superior
mha = MultiheadAttention(embed_dim=8, num_heads=2)
query = torch.rand(4, 1, 8) # (longitud de secuencia, lote, dimensión de embedding)
key = value = query
output, attn_weights = mha(query, key, value, attn_mask=attn_mask, key_padding_mask=key_padding_mask)
De esta forma, la máscara de relleno clave se usa para ignorar los tokens de relleno, y la máscara de atención para controlar la estructura de la atención (por ejemplo, evitar atención a posiciones futuras).