Source code for dual_attention.seq2seq_models
"""
This module implements Encoder-Decoder Sequence-to-Sequence models (both Dual Attention Transformer and standard Transformer).
"""
import torch
from torch import nn
from .transformer_blocks import EncoderBlock, DecoderBlock
from .dual_attn_blocks import DualAttnEncoderBlock, DualAttnDecoderBlock
from .symbol_retrieval import SymbolicAttention, RelationalSymbolicAttention, PositionalSymbolRetriever, PositionRelativeSymbolRetriever
from .positional_encoding import SinusoidalPositionalEncoding, LearnedPositionalEmbeddings
[docs]
class Seq2SeqTransformer(nn.Module):
"""Transformer Language Model"""
[docs]
def __init__(self,
input_spec: dict,
output_spec: dict,
d_model: int,
out_dim: int,
n_layers_enc: int,
n_layers_dec: int,
encoder_kwargs: dict,
decoder_kwargs: dict,
in_block_size: int,
out_block_size: int,
tie_weights: bool = True,
loss_ignore_idx: int = -1):
"""Seq2Seq Encoder-Decoder Transformer.
Parameters
----------
input_spec : dict
description of input format. dictionary with key 'type' with values 'token' or 'vector'.
if 'token', must also have 'vocab_size'. if 'vector', must also have 'dim'.
output_spec : dict
description of output format. dictionary with key 'type' with values 'token' or 'vector'.
if 'token', must also have 'vocab_size'. if 'vector', must also have 'dim'.
d_model : int
model dimension.
out_dim : int
output dimension (e.g., output vocab size)
n_layers_enc : int
number of encoder layers.
n_layers_dec : int
number of decoder layers.
encoder_kwargs : dict
keyword arguments for encoder blocks.
decoder_kwargs : dict
keyword arguments for decoder blocks.
in_block_size : int
block size for input sequence.
out_block_size : int
block size for target sequence.
tie_weights : bool, optional
whether to tie weights between target embedder and final layer weights, by default True
loss_ignore_idx : int, optional
idx of class to ignore when computing loss, by default -1
"""
super().__init__()
self.input_spec = input_spec
self.output_spec = output_spec
self.d_model = d_model
self.out_dim = out_dim
self.n_layers_enc = n_layers_enc
self.n_layers_dec = n_layers_dec
self.encoder_kwargs = encoder_kwargs
self.decoder_kwargs = decoder_kwargs
self.in_block_size = in_block_size
self.out_block_size = out_block_size
self.loss_ignore_idx = loss_ignore_idx
# TODO: make positional embedder configurable (learned or fixed sinusoidal, etc)
if input_spec['type'] == 'token':
source_embedder = torch.nn.Embedding(input_spec['vocab_size'], d_model)
elif input_spec['type'] == 'vector':
source_embedder = torch.nn.Linear(input_spec['dim'], d_model)
else:
raise ValueError(f"input_spec['type'] must be 'token' or 'vector', not {input_spec['type']}")
if output_spec['type'] == 'token':
target_embedder = torch.nn.Embedding(output_spec['vocab_size'], d_model)
elif output_spec['type'] == 'vector':
target_embedder = torch.nn.Linear(output_spec['dim'], d_model)
else:
raise ValueError(f"output_spec['type'] must be 'token' or 'vector', not {output_spec['type']}")
layer_dict = dict(
source_embedder = source_embedder,
target_embedder = target_embedder,
source_pos_embedder = SinusoidalPositionalEncoding(d_model, dropout=0., max_len=in_block_size),
target_pos_embedder = SinusoidalPositionalEncoding(d_model, dropout=0., max_len=out_block_size),
# dropout = nn.Dropout(dropout_rate),
encoder_blocks = nn.ModuleList([EncoderBlock(d_model, **encoder_kwargs) for _ in range(n_layers_enc)]),
decoder_blocks = nn.ModuleList([DecoderBlock(d_model, **decoder_kwargs) for _ in range(n_layers_dec)]),
final_out = nn.Linear(d_model, out_dim)
)
self.layers = nn.ModuleDict(layer_dict)
# weight-tying embedder and final layer
if tie_weights:
self.layers.target_embedder.weights = self.layers.final_out
[docs]
def forward(self, x, y, targets=None):
x = self.layers.source_embedder(x)
y = self.layers.target_embedder(y)
x = self.layers.source_pos_embedder(x)
y = self.layers.target_pos_embedder(y)
for enc_block in self.layers.encoder_blocks:
x = enc_block(x)
for dec_block in self.layers.decoder_blocks:
y = dec_block(y, x)
if targets is not None:
# compute loss if given targets
logits = self.layers.final_out(y)
loss = torch.nn.functional.cross_entropy(logits.view(-1, logits.size(-1)), targets.view(-1),
ignore_index=self.loss_ignore_idx)
else:
logits = self.layers.final_out(y[:, [-1], :])
loss = None
return logits, loss
[docs]
def get_num_params(self):
"""
Return the number of parameters in the model.
"""
n_params = sum(p.numel() for p in self.parameters())
return n_params
[docs]
def estimate_mfu(self, fwdbwd_per_iter, dt):
""" estimate model flops utilization (MFU) in units of A100 bfloat16 peak FLOPS """
# NOTE: Model Flops Utilization (MFU) is a measure of how much of the peak FLOPS of the GPU is being utilized.
# PaLM paper has computed this for standard Transformers
# haven't done this yet for encoder-decoder architectures, so this is a placeholder
return -1.0
[docs]
class Seq2SeqDualAttnTransformer(nn.Module):
"""Dual Attention Transformer Seq2Seq Model"""
[docs]
def __init__(self,
input_spec: dict,
output_spec: dict,
symbol_retrieval: str,
symbol_retrieval_kwargs: dict,
d_model: int,
out_dim: int,
n_layers_enc: int,
n_layers_dec: int,
encoder_kwargs: dict,
decoder_kwargs: dict,
in_block_size: int,
out_block_size: int,
tie_weights: bool = True,
loss_ignore_idx: int = -1):
"""Seq2Seq Encoder-Decoder Dual Attention Transformer
Parameters
----------
input_spec : dict
description of input format. dictionary with key 'type' with values 'token' or 'vector'.
if 'token', must also have 'vocab_size'. if 'vector', must also have 'dim'
output_spec : dict
description of output format. dictionary with key 'type' with values 'token' or 'vector'.
if 'token', must also have 'vocab_size'. if 'vector', must also have 'dim'
symbol_retrieval : str
type of symbol retrieval mechanism. must be one of 'symbolic_attention', 'rel_sym_attn', 'positional_symbols', or 'position_relative'
symbol_retrieval_kwargs : dict
keyword arguments for symbol retrieval mechanism
d_model : int
model dimension
out_dim : int
output dimension (e.g., output vocab size)
n_layers_enc : int
number of encoder layers
n_layers_dec : int
number of decoder layers
encoder_kwargs : dict
keyword arguments for encoder blocks
decoder_kwargs : dict
keyword arguments for decoder blocks
in_block_size : int
block size for input sequence
out_block_size : int
block size for target sequence
tie_weights : bool, optional
whether to tie weights between target embedder and final layer weights, by default True
loss_ignore_idx : int, optional
idx of class to ignore when computing loss, by default -1
"""
super().__init__()
self.input_spec = input_spec
self.output_spec = output_spec
self.d_model = d_model
self.out_dim = out_dim
self.n_layers_enc = n_layers_enc
self.n_layers_dec = n_layers_dec
self.encoder_kwargs = encoder_kwargs
self.decoder_kwargs = decoder_kwargs
self.in_block_size = in_block_size
self.out_block_size = out_block_size
self.loss_ignore_idx = loss_ignore_idx
# TODO: make positional embedder configurable (learned or fixed sinusoidal, etc)
if input_spec['type'] == 'token':
source_embedder = torch.nn.Embedding(input_spec['vocab_size'], d_model)
elif input_spec['type'] == 'vector':
source_embedder = torch.nn.Linear(input_spec['dim'], d_model)
else:
raise ValueError(f"input_spec['type'] must be 'token' or 'vector', not {input_spec['type']}")
# TODO: add option to share embedder between source and target' maybe via output_spec?
if output_spec['type'] == 'token':
target_embedder = torch.nn.Embedding(output_spec['vocab_size'], d_model)
elif output_spec['type'] == 'vector':
target_embedder = torch.nn.Linear(output_spec['dim'], d_model)
else:
raise ValueError(f"output_spec['type'] must be 'token' or 'vector', not {output_spec['type']}")
if symbol_retrieval == 'symbolic_attention':
symbol_retriever = SymbolicAttention(**symbol_retrieval_kwargs)
elif symbol_retrieval == 'rel_sym_attn':
symbol_retriever = RelationalSymbolicAttention(**symbol_retrieval_kwargs)
elif symbol_retrieval == 'positional_symbols':
symbol_retriever = PositionalSymbolRetriever(**symbol_retrieval_kwargs)
elif symbol_retrieval == 'position_relative':
symbol_retriever = PositionRelativeSymbolRetriever(**symbol_retrieval_kwargs)
else:
raise ValueError(f"`symbol_retrieval` must be one of 'symbolic_attention', 'rel_sym_attn', or 'positional_symbols'. received {symbol_retrieval}")
layer_dict = dict(
source_embedder = source_embedder,
target_embedder = target_embedder,
source_pos_embedder = SinusoidalPositionalEncoding(d_model, dropout=0., max_len=in_block_size),
target_pos_embedder = SinusoidalPositionalEncoding(d_model, dropout=0., max_len=out_block_size),
symbol_retriever = symbol_retriever,
# dropout = nn.Dropout(dropout_rate),
encoder_blocks = nn.ModuleList([DualAttnEncoderBlock(d_model, **encoder_kwargs) for _ in range(n_layers_enc)]),
decoder_blocks = nn.ModuleList([DualAttnDecoderBlock(d_model, **decoder_kwargs) for _ in range(n_layers_enc)]),
final_out = nn.Linear(d_model, out_dim)
)
self.layers = nn.ModuleDict(layer_dict)
# weight-tying embedder and final layer
if tie_weights:
self.layers.target_embedder.weights = self.layers.final_out
[docs]
def forward(self, x, y, targets=None):
x = self.layers.source_embedder(x)
y = self.layers.target_embedder(y)
x = self.layers.source_pos_embedder(x)
y = self.layers.target_pos_embedder(y)
for enc_block in self.layers.encoder_blocks:
symbols = self.layers.symbol_retriever(x)
x = enc_block(x, symbols)
for dec_block in self.layers.decoder_blocks:
symbols = self.layers.symbol_retriever(y)
y = dec_block(y, x, symbols=symbols)
if targets is not None:
# compute loss if given targets
logits = self.layers.final_out(y)
loss = torch.nn.functional.cross_entropy(
logits.view(-1, logits.size(-1)), targets.view(-1), ignore_index=self.loss_ignore_idx)
else:
logits = self.layers.final_out(y[:, [-1], :])
loss = None
return logits, loss
[docs]
def get_num_params(self):
"""
Return the number of parameters in the model.
"""
n_params = sum(p.numel() for p in self.parameters())
return n_params
[docs]
def estimate_mfu(self, fwdbwd_per_iter, dt):
""" estimate model flops utilization (MFU) in units of A100 bfloat16 peak FLOPS """
# NOTE: Model Flops Utilization (MFU) is a measure of how much of the peak FLOPS of the GPU is being utilized.
# PaLM paper has computed this for standard Transformers
# haven't done this yet for encoder-decoder architectures, so this is a placeholder
return -1.0
[docs]
def configure_optimizers(model, weight_decay, learning_rate, betas, device_type):
# start with all of the candidate parameters
param_dict = {pn: p for pn, p in model.named_parameters()}
# filter out those that do not require grad
param_dict = {pn: p for pn, p in param_dict.items() if p.requires_grad}
# create optim groups. Any parameters that is 2D will be weight decayed, otherwise no.
# i.e. all weight tensors in matmuls + embeddings decay, all biases and layernorms don't.
decay_params = [p for n, p in param_dict.items() if p.dim() >= 2]
nodecay_params = [p for n, p in param_dict.items() if p.dim() < 2]
optim_groups = [
{'params': decay_params, 'weight_decay': weight_decay},
{'params': nodecay_params, 'weight_decay': 0.0}
]
num_decay_params = sum(p.numel() for p in decay_params)
num_nodecay_params = sum(p.numel() for p in nodecay_params)
print(f"num decayed parameter tensors: {len(decay_params)}, with {num_decay_params:,} parameters")
print(f"num non-decayed parameter tensors: {len(nodecay_params)}, with {num_nodecay_params:,} parameters")
# Create AdamW optimizer and use the fused version if it is available
use_fused = (device_type == 'cuda')
optimizer = torch.optim.AdamW(optim_groups, lr=learning_rate, betas=betas, fused=use_fused)
print(f"using fused AdamW: {use_fused}")
return optimizer