"""
This module implements forward calls for the Dual Attention Transformer (DAT) model
that return intermediate results for visualization purposes.
"""
import torch
[docs]
def symbolic_attn_forward_get_weights(mod, x):
'''a variant of the forward call for symbolic attention that returns the attention weights'''
mod.eval()
batch_size, seq_len, dim = x.size()
# create query from input
query = mod.q_proj(x)
query = query.view(batch_size, seq_len, mod.n_heads, dim // mod.n_heads).transpose(1, 2)
# create keys from template features
key = mod.template_features.view(mod.n_symbols, mod.n_heads, mod.d_model // mod.n_heads).transpose(0, 1)
key = mod._repeat_kv(key, batch_size)
# create values from symbol library
value = mod.symbol_library.view(mod.n_symbols, mod.n_heads, mod.d_model // mod.n_heads).transpose(0, 1)
value = mod._repeat_kv(value, batch_size)
# retrieved_symbols = torch.nn.functional.scaled_dot_product_attention(
# query, key, value,
# scale=self.scale, dropout_p=self.dropout, attn_mask=None, is_causal=False)
scale = mod.scale if mod.scale is not None else (mod.d_model/mod.n_heads) ** -0.5
attn_scores = torch.matmul(query, key.transpose(2, 3)) * scale
attn_scores = torch.nn.functional.softmax(attn_scores, dim=-1)
retrieved_symbols = torch.matmul(attn_scores, value)
retrieved_symbols = retrieved_symbols.transpose(1, 2).contiguous().view(batch_size, seq_len, dim)
return retrieved_symbols, attn_scores
[docs]
def block_forward_get_weights(mod, x, symbols, freqs_cos=None, freqs_sin=None):
'''a variant of the forward call for a block that returns the attention weights'''
mod.eval()
def _compute_dual_attn(mod, x, symbols, freqs_cos=None, freqs_sin=None):
x, sa_attn_scores, (ra_attn_scores, ra_rels) = mod.dual_attn(x, symbols,
need_weights=True, is_causal=mod.causal,
freqs_cos=freqs_cos, freqs_sin=freqs_sin)
x = mod.dropout(x) # dropout
return x, sa_attn_scores, ra_attn_scores, ra_rels
if mod.norm_first:
attn_out, sa_attn_scores, ra_attn_scores, ra_rels = _compute_dual_attn(mod, mod.norm1(x), symbols, freqs_cos=freqs_cos, freqs_sin=freqs_sin)
x = x + attn_out
x = x + mod._apply_ff_block(mod.norm2(x))
else:
attn_out, sa_attn_scores, ra_attn_scores = _compute_dual_attn(mod, x, symbols, freqs_cos=freqs_cos, freqs_sin=freqs_sin)
x = mod.norm1(x + attn_out)
x = mod.dropout(x)
x = mod.norm2(x + mod._apply_ff_block(x))
return x, sa_attn_scores, ra_attn_scores, ra_rels