torch.nn.attention.bias.CausalBias
-
class torch.nn.attention.bias.CausalBias(variant, seq_len_q, seq_len_kv)[source] -
Смещение, представляющее причинные шаблоны внимания. Обзор структуры смещения см. в перечислении
CausalVariant.Этот класс используется для определения причинных (треугольных) смещений внимания. Для создания смещения предусмотрены две фабричные функции:
causal_upper_left()иcausal_lower_right().Пример:
from torch.nn.attention.bias import causal_lower_right bsz, num_heads, seqlen_q, seqlen_kv, head_dim = 32, 8, 4, 12, 8 # Create a lower-right causal bias attn_bias = causal_lower_right(seqlen_q, seqlen_kv) q = torch.randn( bsz, num_heads, seqlen_q, head_dim, device="cuda", dtype=torch.float16 ) k = torch.randn( bsz, num_heads, seqlen_kv, head_dim, device="cuda", dtype=torch.float16 ) v = torch.randn( bsz, num_heads, seqlen_kv, head_dim, device="cuda", dtype=torch.float16 ) out = F.scaled_dot_product_attention(q, k, v, attn_bias)Предупреждение
Этот класс является прототипом и может измениться.
© 2026, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://docs.pytorch.org/docs/2.14/generated/torch.nn.attention.bias.CausalBias.html