Compare commits
3
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
29df4df121 | ||
|
|
260e0c691d | ||
|
|
a1873aef57 |
Executable
+1
@@ -0,0 +1 @@
|
||||
python benchmarks.py --models bart-large-cnn --batch_sizes 2 --torch
|
||||
@@ -0,0 +1,67 @@
|
||||
from transformers import *
|
||||
import torch
|
||||
DEFAULT_DEVICE = 'cuda' if torch.cuda.is_available() else 'cpu'
|
||||
def runner(source_path, out_file, batch_size=8, device=DEFAULT_DEVICE, prof_generate=False):
|
||||
|
||||
tokenizer = BartTokenizer.from_pretrained('bart-large')
|
||||
lns = [" " + x.rstrip() for x in open(source_path).readlines()][:batch_size]
|
||||
|
||||
dct = tokenizer.batch_encode_plus(lns, max_length=1024, return_tensors="pt", pad_to_max_length=True)
|
||||
ids = dct['input_ids'].to(DEFAULT_DEVICE)
|
||||
msk = dct['attention_mask'].to(DEFAULT_DEVICE)
|
||||
model = BartForConditionalGeneration.from_pretrained('bart-large-cnn', output_past=prof_generate).to(DEFAULT_DEVICE)
|
||||
model.log_mem('starting')
|
||||
if prof_generate:
|
||||
|
||||
summaries = model.generate(
|
||||
input_ids=ids,
|
||||
attention_mask=msk,
|
||||
num_beams=4,
|
||||
length_penalty=2.0,
|
||||
max_length=140 + 2, # +2 from original because we start at step=1 and stop before max_length
|
||||
min_length=55 + 1, # +1 from original because we start at step=1
|
||||
no_repeat_ngram_size=3,
|
||||
early_stopping=True,
|
||||
do_sample=False,
|
||||
decoder_start_token_id=model.config.eos_token_ids[0],
|
||||
)
|
||||
model.log_mem('done')
|
||||
dec = [tokenizer.decode(s) for s in summaries]
|
||||
print(dec[0])
|
||||
else:
|
||||
#model.decoder.generation_mode = Fals
|
||||
with torch.no_grad():
|
||||
model(
|
||||
input_ids=ids,
|
||||
attention_mask=msk,
|
||||
)
|
||||
|
||||
log_df = model.combine_logs()
|
||||
log_df.to_csv(out_file)
|
||||
|
||||
|
||||
import argparse
|
||||
if __name__ == '__main__':
|
||||
parser = argparse.ArgumentParser()
|
||||
|
||||
parser.add_argument(
|
||||
"output_path", type=str, help="where to save summaries",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--source_path", type=str, default="/home/shleifer/transformers_fork/notebooks/test.source",
|
||||
help="like cnn_dm/test.source", required=False
|
||||
)
|
||||
parser.add_argument(
|
||||
"--device", type=str, required=False, default=DEFAULT_DEVICE, help="cuda, cuda:1, cpu etc.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--bs", type=int, default=8, required=False, help="batch size: how many to summarize at a time",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--do-generate", action='store_true', required=False, help="batch size: how many to summarize at a time",
|
||||
)
|
||||
args = parser.parse_args()
|
||||
runner(args.source_path, args.output_path, batch_size=args.bs, device=args.device, prof_generate=args.do_generate)
|
||||
|
||||
|
||||
|
||||
@@ -19,10 +19,10 @@ from .configuration_utils import PretrainedConfig
|
||||
|
||||
|
||||
ALBERT_PRETRAINED_CONFIG_ARCHIVE_MAP = {
|
||||
"albert-base-v1": "https://s3.amazonaws.com/models.huggingface.co/bert/albert-base-v1-config.json",
|
||||
"albert-large-v1": "https://s3.amazonaws.com/models.huggingface.co/bert/albert-large-v1-config.json",
|
||||
"albert-xlarge-v1": "https://s3.amazonaws.com/models.huggingface.co/bert/albert-xlarge-v1-config.json",
|
||||
"albert-xxlarge-v1": "https://s3.amazonaws.com/models.huggingface.co/bert/albert-xxlarge-v1-config.json",
|
||||
"albert-base-v1": "https://s3.amazonaws.com/models.huggingface.co/bert/albert-base-config.json",
|
||||
"albert-large-v1": "https://s3.amazonaws.com/models.huggingface.co/bert/albert-large-config.json",
|
||||
"albert-xlarge-v1": "https://s3.amazonaws.com/models.huggingface.co/bert/albert-xlarge-config.json",
|
||||
"albert-xxlarge-v1": "https://s3.amazonaws.com/models.huggingface.co/bert/albert-xxlarge-config.json",
|
||||
"albert-base-v2": "https://s3.amazonaws.com/models.huggingface.co/bert/albert-base-v2-config.json",
|
||||
"albert-large-v2": "https://s3.amazonaws.com/models.huggingface.co/bert/albert-large-v2-config.json",
|
||||
"albert-xlarge-v2": "https://s3.amazonaws.com/models.huggingface.co/bert/albert-xlarge-v2-config.json",
|
||||
|
||||
@@ -33,10 +33,10 @@ logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
ALBERT_PRETRAINED_MODEL_ARCHIVE_MAP = {
|
||||
"albert-base-v1": "https://s3.amazonaws.com/models.huggingface.co/bert/albert-base-v1-pytorch_model.bin",
|
||||
"albert-large-v1": "https://s3.amazonaws.com/models.huggingface.co/bert/albert-large-v1-pytorch_model.bin",
|
||||
"albert-xlarge-v1": "https://s3.amazonaws.com/models.huggingface.co/bert/albert-xlarge-v1-pytorch_model.bin",
|
||||
"albert-xxlarge-v1": "https://s3.amazonaws.com/models.huggingface.co/bert/albert-xxlarge-v1-pytorch_model.bin",
|
||||
"albert-base-v1": "https://s3.amazonaws.com/models.huggingface.co/bert/albert-base-pytorch_model.bin",
|
||||
"albert-large-v1": "https://s3.amazonaws.com/models.huggingface.co/bert/albert-large-pytorch_model.bin",
|
||||
"albert-xlarge-v1": "https://s3.amazonaws.com/models.huggingface.co/bert/albert-xlarge-pytorch_model.bin",
|
||||
"albert-xxlarge-v1": "https://s3.amazonaws.com/models.huggingface.co/bert/albert-xxlarge-pytorch_model.bin",
|
||||
"albert-base-v2": "https://s3.amazonaws.com/models.huggingface.co/bert/albert-base-v2-pytorch_model.bin",
|
||||
"albert-large-v2": "https://s3.amazonaws.com/models.huggingface.co/bert/albert-large-v2-pytorch_model.bin",
|
||||
"albert-xlarge-v2": "https://s3.amazonaws.com/models.huggingface.co/bert/albert-xlarge-v2-pytorch_model.bin",
|
||||
|
||||
@@ -25,7 +25,8 @@ from .activations import ACT2FN
|
||||
from .configuration_bart import BartConfig
|
||||
from .file_utils import add_start_docstrings, add_start_docstrings_to_callable
|
||||
from .modeling_utils import PreTrainedModel, create_position_ids_from_input_ids
|
||||
|
||||
from durbango.logging_utils import LoggingMixin
|
||||
from durbango.torch_utils import print_tensor_sizes, local_sizeof, get_tensor_shapes_and_pointers
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -109,7 +110,7 @@ def _prepare_bart_decoder_inputs(
|
||||
return decoder_input_ids, decoder_attn_mask
|
||||
|
||||
|
||||
class PretrainedBartModel(PreTrainedModel):
|
||||
class PretrainedBartModel(PreTrainedModel, LoggingMixin):
|
||||
config_class = BartConfig
|
||||
base_model_prefix = "model"
|
||||
pretrained_model_archive_map = BART_PRETRAINED_MODEL_ARCHIVE_MAP
|
||||
@@ -185,9 +186,10 @@ def make_padding_mask(input_ids, padding_idx=1):
|
||||
|
||||
|
||||
# Helper Modules
|
||||
from durbango.torch_utils import get_shapes
|
||||
|
||||
|
||||
class EncoderLayer(nn.Module):
|
||||
class EncoderLayer(nn.Module, LoggingMixin):
|
||||
def __init__(self, config: BartConfig):
|
||||
super().__init__()
|
||||
self.embed_dim = config.d_model
|
||||
@@ -216,7 +218,8 @@ class EncoderLayer(nn.Module):
|
||||
encoded output of shape `(seq_len, batch, embed_dim)`
|
||||
"""
|
||||
residual = x
|
||||
x, attn_weights = self.self_attn(query=x, key=x, key_padding_mask=encoder_padding_mask,)
|
||||
x, attn_weights = self.self_attn(query=x, key=x, key_padding_mask=encoder_padding_mask, update_layer_state=False,)
|
||||
|
||||
x = F.dropout(x, p=self.dropout, training=self.training)
|
||||
x = residual + x
|
||||
x = self.self_attn_layer_norm(x)
|
||||
@@ -230,8 +233,8 @@ class EncoderLayer(nn.Module):
|
||||
x = self.final_layer_norm(x)
|
||||
return x, attn_weights
|
||||
|
||||
|
||||
class BartEncoder(nn.Module):
|
||||
import gc
|
||||
class BartEncoder(nn.Module, LoggingMixin):
|
||||
"""
|
||||
Transformer encoder consisting of *config.encoder_layers* self attention layers. Each layer
|
||||
is a :class:`EncoderLayer`.
|
||||
@@ -282,16 +285,19 @@ class BartEncoder(nn.Module):
|
||||
attention_mask = attention_mask.eq(0)
|
||||
|
||||
inputs_embeds = self.embed_tokens(input_ids)
|
||||
embed_pos = self.embed_positions(input_ids)
|
||||
x = inputs_embeds + embed_pos
|
||||
x = inputs_embeds + self.embed_positions(input_ids)
|
||||
x = self.layernorm_embedding(x)
|
||||
x = F.dropout(x, p=self.dropout, training=self.training)
|
||||
assert not (self.output_attentions or self.output_hidden_states)
|
||||
|
||||
# B x T x C -> T x B x C
|
||||
x = x.transpose(0, 1)
|
||||
|
||||
self.log_mem('encoder: starting_loop')
|
||||
encoder_states, all_attentions = [], []
|
||||
for encoder_layer in self.layers:
|
||||
#rdd_start = print_tensor_sizes()
|
||||
#rdd_start.to_csv(f'rdd_start.csv')
|
||||
for i, encoder_layer in enumerate(self.layers):
|
||||
|
||||
if self.output_hidden_states:
|
||||
encoder_states.append(x)
|
||||
# add LayerDrop (see https://arxiv.org/abs/1909.11556 for description)
|
||||
@@ -299,19 +305,16 @@ class BartEncoder(nn.Module):
|
||||
if self.training and (dropout_probability < self.layerdrop): # skip the layer
|
||||
attn = None
|
||||
else:
|
||||
x, attn = encoder_layer(x, attention_mask)
|
||||
|
||||
if self.output_attentions:
|
||||
all_attentions.append(attn)
|
||||
|
||||
if self.output_hidden_states:
|
||||
encoder_states.append(x)
|
||||
|
||||
x, _ = encoder_layer(x, attention_mask)
|
||||
assert len(encoder_states) == 0
|
||||
assert len(all_attentions) == 0
|
||||
#self.log_mem(f'x: {x.shape}, attn: {attn.shape}')
|
||||
self.log_mem(f'Encoder: called layer {i}')
|
||||
encoder_states = [hidden_state.transpose(0, 1) for hidden_state in encoder_states]
|
||||
return x, encoder_states, all_attentions
|
||||
|
||||
|
||||
class DecoderLayer(nn.Module):
|
||||
class DecoderLayer(nn.Module, LoggingMixin):
|
||||
def __init__(self, config: BartConfig):
|
||||
super().__init__()
|
||||
self.embed_dim = config.d_model
|
||||
@@ -374,7 +377,7 @@ class DecoderLayer(nn.Module):
|
||||
) # just self_attn weights for now, following t5, layer_state = cache for decoding
|
||||
|
||||
|
||||
class BartDecoder(nn.Module):
|
||||
class BartDecoder(nn.Module, LoggingMixin):
|
||||
"""
|
||||
Transformer decoder consisting of *config.decoder_layers* layers. Each layer
|
||||
is a :class:`DecoderLayer`.
|
||||
@@ -435,6 +438,8 @@ class BartDecoder(nn.Module):
|
||||
encoder_padding_mask = encoder_padding_mask.eq(0)
|
||||
|
||||
# embed positions
|
||||
|
||||
self.log_mem('decoder: embedded positions')
|
||||
positions = self.embed_positions(input_ids, generation_mode=generation_mode)
|
||||
|
||||
if generation_mode:
|
||||
@@ -443,6 +448,7 @@ class BartDecoder(nn.Module):
|
||||
assert input_ids.ne(self.padding_idx).any()
|
||||
|
||||
x = self.embed_tokens(input_ids)
|
||||
self.log_mem('decoder: embedded tokens')
|
||||
x += positions
|
||||
|
||||
x = self.layernorm_embedding(x)
|
||||
@@ -464,6 +470,7 @@ class BartDecoder(nn.Module):
|
||||
x, layer_self_attn, layer_past = decoder_layer(
|
||||
x, encoder_hidden_states, encoder_padding_mask, layer_state=layer_state, attention_mask=combined_mask,
|
||||
)
|
||||
self.log_mem(f'decoder: called attn {i}')
|
||||
|
||||
if self.output_past:
|
||||
next_decoder_cache.append(layer_past.copy())
|
||||
@@ -483,6 +490,7 @@ class BartDecoder(nn.Module):
|
||||
return x, next_cache, all_hidden_states, list(all_self_attns)
|
||||
|
||||
|
||||
|
||||
def _reorder_buffer(attn_cache, new_order):
|
||||
for k, input_buffer_k in attn_cache.items():
|
||||
if input_buffer_k is not None:
|
||||
@@ -490,8 +498,8 @@ def _reorder_buffer(attn_cache, new_order):
|
||||
return attn_cache
|
||||
|
||||
|
||||
class SelfAttention(nn.Module):
|
||||
"""Multi-headed attention from 'Attention Is All You Need' paper"""
|
||||
class SelfAttention(nn.Module, LoggingMixin):
|
||||
"""Multi-headed attention from "Attention Is All You Need"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
@@ -519,11 +527,16 @@ class SelfAttention(nn.Module):
|
||||
def _shape(self, tensor, dim_0, bsz):
|
||||
return tensor.contiguous().view(dim_0, bsz * self.num_heads, self.head_dim).transpose(0, 1)
|
||||
|
||||
|
||||
def log_mem(self, msg='', verbose=False):
|
||||
super().log_mem(msg=f'{self.cache_key}_attn:{msg}', verbose=verbose)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
query,
|
||||
key: Optional[Tensor],
|
||||
key_padding_mask: Optional[Tensor] = None,
|
||||
update_layer_state=True,
|
||||
layer_state: Optional[Dict[str, Optional[Tensor]]] = None,
|
||||
attn_mask: Optional[Tensor] = None,
|
||||
) -> Tuple[Tensor, Optional[Tensor]]:
|
||||
@@ -544,6 +557,7 @@ class SelfAttention(nn.Module):
|
||||
layer_state = {}
|
||||
|
||||
q = self.q_proj(query) * self.scaling
|
||||
self.log_mem('\tq_proj')
|
||||
if static_kv:
|
||||
if key is None:
|
||||
k = v = None
|
||||
@@ -554,29 +568,39 @@ class SelfAttention(nn.Module):
|
||||
k = self.k_proj(query)
|
||||
v = self.v_proj(query)
|
||||
|
||||
|
||||
q = self._shape(q, tgt_len, bsz)
|
||||
self.log_mem(f'\tq_reshape -> {q.shape}')
|
||||
if k is not None:
|
||||
k = self._shape(k, -1, bsz)
|
||||
self.log_mem(f'\t done reshaping k,v ->, {k.shape}')
|
||||
if v is not None:
|
||||
v = self._shape(v, -1, bsz)
|
||||
|
||||
|
||||
if saved_state is not None:
|
||||
self.log_mem('\t about to use saved_state')
|
||||
k, v, key_padding_mask = self._use_saved_state(k, v, saved_state, key_padding_mask, static_kv, bsz)
|
||||
|
||||
# Update cache
|
||||
layer_state[self.cache_key] = {
|
||||
"prev_key": k.view(bsz, self.num_heads, -1, self.head_dim),
|
||||
"prev_value": v.view(bsz, self.num_heads, -1, self.head_dim),
|
||||
"prev_key_padding_mask": key_padding_mask if not static_kv else None,
|
||||
}
|
||||
if update_layer_state:
|
||||
layer_state[self.cache_key] = {
|
||||
"prev_key": k.view(bsz, self.num_heads, -1, self.head_dim),
|
||||
"prev_value": v.view(bsz, self.num_heads, -1, self.head_dim),
|
||||
"prev_key_padding_mask": key_padding_mask if not static_kv else None,
|
||||
}
|
||||
self.log_mem('\t attn: done layer_state')
|
||||
|
||||
assert k is not None
|
||||
src_len = k.size(1)
|
||||
self.log_mem('\t attn: before BMM(q,k)')
|
||||
attn_weights = torch.bmm(q, k.transpose(1, 2))
|
||||
self.log_mem('\t attn: done BMM(q,k)')
|
||||
assert attn_weights.size() == (bsz * self.num_heads, tgt_len, src_len)
|
||||
|
||||
if attn_mask is not None:
|
||||
attn_weights = attn_weights.view(bsz, self.num_heads, tgt_len, src_len) + attn_mask
|
||||
self.log_mem('\t attn: done causal mask')
|
||||
attn_weights = attn_weights.view(bsz * self.num_heads, tgt_len, src_len)
|
||||
|
||||
# This is part of a workaround to get around fork/join parallelism not supporting Optional types.
|
||||
@@ -584,21 +608,26 @@ class SelfAttention(nn.Module):
|
||||
key_padding_mask = None
|
||||
assert key_padding_mask is None or key_padding_mask.size()[:2] == (bsz, src_len,)
|
||||
|
||||
if key_padding_mask is not None: # don't attend to padding symbols
|
||||
if key_padding_mask is not None: # shape (bsz, src_len)
|
||||
attn_weights = attn_weights.view(bsz, self.num_heads, tgt_len, src_len)
|
||||
reshaped = key_padding_mask.unsqueeze(1).unsqueeze(2)
|
||||
attn_weights = attn_weights.masked_fill(reshaped, float("-inf"))
|
||||
attn_weights = attn_weights.masked_fill(key_padding_mask.unsqueeze(1).unsqueeze(2), float("-inf"))
|
||||
self.log_mem('\t attn: done masked_fill')
|
||||
attn_weights = attn_weights.view(bsz * self.num_heads, tgt_len, src_len)
|
||||
attn_weights = F.softmax(attn_weights, dim=-1)
|
||||
self.log_mem('\t attn: done softmax')
|
||||
attn_probs = F.dropout(attn_weights, p=self.dropout, training=self.training,)
|
||||
|
||||
|
||||
assert v is not None
|
||||
attn_output = torch.bmm(attn_probs, v)
|
||||
self.log_mem('\t attn: done BMM(probs, v)')
|
||||
assert attn_output.size() == (bsz * self.num_heads, tgt_len, self.head_dim)
|
||||
attn_output = attn_output.transpose(0, 1).contiguous().view(tgt_len, bsz, embed_dim)
|
||||
self.log_mem('\t attn: done view(output)')
|
||||
attn_output = self.out_proj(attn_output)
|
||||
attn_weights = attn_weights.view(bsz, self.num_heads, tgt_len, src_len)
|
||||
return attn_output, attn_weights
|
||||
self.log_mem('\t attn: done out_proj')
|
||||
#attn_weights = attn_weights.view(bsz, self.num_heads, tgt_len, src_len)
|
||||
return attn_output, None
|
||||
|
||||
def _use_saved_state(self, k, v, saved_state, key_padding_mask, static_kv, bsz):
|
||||
# saved states are stored with shape (bsz, num_heads, seq_len, head_dim)
|
||||
@@ -655,7 +684,7 @@ class SelfAttention(nn.Module):
|
||||
return new_key_padding_mask
|
||||
|
||||
|
||||
class BartClassificationHead(nn.Module):
|
||||
class BartClassificationHead(nn.Module, LoggingMixin):
|
||||
"""Head for sentence-level classification tasks."""
|
||||
|
||||
# This can trivially be shared with RobertaClassificationHead
|
||||
@@ -727,6 +756,8 @@ def _filter_out_falsey_values(tup) -> Tuple:
|
||||
|
||||
# Public API
|
||||
|
||||
import time
|
||||
import pandas as pd
|
||||
|
||||
@add_start_docstrings(
|
||||
"The bare BART Model outputting raw hidden-states without any specific head on top.", BART_START_DOCSTRING,
|
||||
@@ -743,6 +774,7 @@ class BartModel(PretrainedBartModel):
|
||||
self.encoder = BartEncoder(config, self.shared)
|
||||
self.decoder = BartDecoder(config, self.shared)
|
||||
|
||||
|
||||
self.init_weights()
|
||||
|
||||
@add_start_docstrings_to_callable(BART_INPUTS_DOCSTRING)
|
||||
@@ -758,6 +790,9 @@ class BartModel(PretrainedBartModel):
|
||||
):
|
||||
|
||||
# make masks if user doesn't supply
|
||||
if encoder_outputs is None:
|
||||
encoder_outputs = self.encoder(input_ids=input_ids, attention_mask=attention_mask)
|
||||
assert isinstance(encoder_outputs, tuple)
|
||||
if not generation_mode:
|
||||
decoder_input_ids, decoder_attention_mask = _prepare_bart_decoder_inputs(
|
||||
self.config,
|
||||
@@ -767,9 +802,6 @@ class BartModel(PretrainedBartModel):
|
||||
mask_dtype=self.shared.weight.dtype,
|
||||
)
|
||||
assert decoder_input_ids is not None
|
||||
if encoder_outputs is None:
|
||||
encoder_outputs = self.encoder(input_ids=input_ids, attention_mask=attention_mask)
|
||||
assert isinstance(encoder_outputs, tuple)
|
||||
# decoder outputs consists of (dec_features, layer_state, dec_hidden, dec_attn)
|
||||
decoder_outputs = self.decoder(
|
||||
decoder_input_ids,
|
||||
@@ -804,10 +836,10 @@ class BartForConditionalGeneration(PretrainedBartModel):
|
||||
|
||||
def __init__(self, config: BartConfig):
|
||||
super().__init__(config)
|
||||
# if base_model is None:
|
||||
base_model = BartModel(config)
|
||||
self.model = base_model
|
||||
self.lm_head = _make_linear_from_emb(self.model.shared)
|
||||
# if base_model is Nones:
|
||||
#self.log_mem('pre-init')
|
||||
self.model = BartModel(config)
|
||||
#self.lm_head = _make_linear_from_emb(self.model.shared)
|
||||
|
||||
def tie_weights(self):
|
||||
pass # hack to prevent changing lm_head.out_features. The input and output embeddings are still the same.
|
||||
@@ -866,6 +898,7 @@ class BartForConditionalGeneration(PretrainedBartModel):
|
||||
tokenizer.decode(predictions).split()
|
||||
# ['good', 'great', 'all', 'really', 'very']
|
||||
"""
|
||||
self.model.log_mem('before BartModel.forward')
|
||||
outputs = self.model(
|
||||
input_ids,
|
||||
attention_mask=attention_mask,
|
||||
@@ -875,7 +908,10 @@ class BartForConditionalGeneration(PretrainedBartModel):
|
||||
decoder_cached_states=decoder_cached_states,
|
||||
generation_mode=generation_mode,
|
||||
)
|
||||
lm_logits = self.lm_head(outputs[0])
|
||||
self.model.log_mem('after call, before lm_head')
|
||||
lm_logits = F.linear(outputs[0], self.model.shared.weight)
|
||||
#lm_logits = self.lm_head(outputs[0])
|
||||
self.model.log_mem('after lm_head')
|
||||
outputs = (lm_logits,) + outputs[1:] # Add hidden states and attention if they are here
|
||||
if lm_labels is not None:
|
||||
loss_fct = nn.CrossEntropyLoss()
|
||||
@@ -885,6 +921,7 @@ class BartForConditionalGeneration(PretrainedBartModel):
|
||||
|
||||
return outputs
|
||||
|
||||
|
||||
def prepare_inputs_for_generation(self, decoder_input_ids, past, attention_mask, **kwargs):
|
||||
assert past is not None, "past has to be defined for encoder_outputs"
|
||||
|
||||
@@ -893,7 +930,8 @@ class BartForConditionalGeneration(PretrainedBartModel):
|
||||
encoder_outputs, decoder_cached_states = past, None
|
||||
else:
|
||||
encoder_outputs, decoder_cached_states = past
|
||||
|
||||
self.log_mem(f'encoder_outputs.shape: {encoder_outputs[0].shape}')
|
||||
self.log_mem(f'decoder_input_ids.shape: {decoder_input_ids.shape}')
|
||||
return {
|
||||
"input_ids": None, # encoder_outputs is defined. input_ids not needed
|
||||
"encoder_outputs": encoder_outputs,
|
||||
@@ -932,7 +970,7 @@ class BartForConditionalGeneration(PretrainedBartModel):
|
||||
return self.model.encoder
|
||||
|
||||
def get_output_embeddings(self):
|
||||
return self.lm_head
|
||||
return _make_linear_from_emb(self.model.shared) # make it on the fly
|
||||
|
||||
|
||||
@add_start_docstrings(
|
||||
|
||||
@@ -814,9 +814,10 @@ class T5Model(T5PreTrainedModel):
|
||||
|
||||
return decoder_outputs + encoder_outputs
|
||||
|
||||
from durbango.logging_utils import LoggingMixin
|
||||
|
||||
@add_start_docstrings("""T5 Model with a `language modeling` head on top. """, T5_START_DOCSTRING, T5_INPUTS_DOCSTRING)
|
||||
class T5ForConditionalGeneration(T5PreTrainedModel):
|
||||
class T5ForConditionalGeneration(T5PreTrainedModel, LoggingMixin):
|
||||
r"""
|
||||
**lm_labels**: (`optional`) ``torch.LongTensor`` of shape ``(batch_size, sequence_length)``:
|
||||
Labels for computing the masked language modeling loss.
|
||||
|
||||
@@ -896,6 +896,17 @@ class PreTrainedModel(nn.Module, ModuleUtilsMixin):
|
||||
effective_batch_size = batch_size
|
||||
effective_batch_mult = 1
|
||||
|
||||
if self.config.is_encoder_decoder:
|
||||
assert bos_token_id is not None, "Encoder Decoder Models need to have a bos_token_id"
|
||||
assert hasattr(self, "get_encoder"), "{} should have a 'get_encoder' function defined".format(self)
|
||||
assert callable(self.get_encoder), "{} should be a method".format(self.get_encoder)
|
||||
|
||||
# get encoder and store encoder outputs
|
||||
encoder = self.get_encoder()
|
||||
|
||||
encoder_outputs = encoder(input_ids, attention_mask=attention_mask)
|
||||
assert encoder_outputs[0].shape[1] == input_ids.shape[0]
|
||||
|
||||
# Expand input ids if num_beams > 1 or num_return_sequences > 1
|
||||
if num_return_sequences > 1 or num_beams > 1:
|
||||
input_ids_len = input_ids.shape[-1]
|
||||
@@ -911,16 +922,8 @@ class PreTrainedModel(nn.Module, ModuleUtilsMixin):
|
||||
effective_batch_size * num_beams, input_ids_len
|
||||
) # shape: (batch_size * num_return_sequences * num_beams, cur_len)
|
||||
|
||||
|
||||
if self.config.is_encoder_decoder:
|
||||
assert bos_token_id is not None, "Encoder Decoder Models need to have a bos_token_id"
|
||||
assert hasattr(self, "get_encoder"), "{} should have a 'get_encoder' function defined".format(self)
|
||||
assert callable(self.get_encoder), "{} should be a method".format(self.get_encoder)
|
||||
|
||||
# get encoder and store encoder outputs
|
||||
encoder = self.get_encoder()
|
||||
|
||||
encoder_outputs = encoder(input_ids, attention_mask=attention_mask)
|
||||
|
||||
# create empty decoder_input_ids
|
||||
input_ids = torch.full(
|
||||
(effective_batch_size * num_beams, 1),
|
||||
@@ -929,10 +932,16 @@ class PreTrainedModel(nn.Module, ModuleUtilsMixin):
|
||||
device=next(self.parameters()).device,
|
||||
)
|
||||
cur_len = 1
|
||||
self.log_mem('about to rearrange')
|
||||
assert batch_size == encoder_outputs[0].shape[1], "NEED MSG"
|
||||
expanded_index = torch.arange(batch_size).view(-1, 1).repeat(1, num_beams* effective_batch_mult).view(-1).to(input_ids.device)
|
||||
encoder_outputs = (encoder_outputs[0].index_select(1, expanded_index), *encoder_outputs[1:])
|
||||
self.log_mem('done rearrange')
|
||||
else:
|
||||
encoder_outputs = None
|
||||
cur_len = input_ids.shape[-1]
|
||||
|
||||
|
||||
if num_beams > 1:
|
||||
output = self._generate_beam_search(
|
||||
input_ids,
|
||||
|
||||
@@ -29,10 +29,10 @@ VOCAB_FILES_NAMES = {"vocab_file": "spiece.model"}
|
||||
|
||||
PRETRAINED_VOCAB_FILES_MAP = {
|
||||
"vocab_file": {
|
||||
"albert-base-v1": "https://s3.amazonaws.com/models.huggingface.co/bert/albert-base-v1-spiece.model",
|
||||
"albert-large-v1": "https://s3.amazonaws.com/models.huggingface.co/bert/albert-large-v1-spiece.model",
|
||||
"albert-xlarge-v1": "https://s3.amazonaws.com/models.huggingface.co/bert/albert-xlarge-v1-spiece.model",
|
||||
"albert-xxlarge-v1": "https://s3.amazonaws.com/models.huggingface.co/bert/albert-xxlarge-v1-spiece.model",
|
||||
"albert-base-v1": "https://s3.amazonaws.com/models.huggingface.co/bert/albert-base-spiece.model",
|
||||
"albert-large-v1": "https://s3.amazonaws.com/models.huggingface.co/bert/albert-large-spiece.model",
|
||||
"albert-xlarge-v1": "https://s3.amazonaws.com/models.huggingface.co/bert/albert-xlarge-spiece.model",
|
||||
"albert-xxlarge-v1": "https://s3.amazonaws.com/models.huggingface.co/bert/albert-xxlarge-spiece.model",
|
||||
"albert-base-v2": "https://s3.amazonaws.com/models.huggingface.co/bert/albert-base-v2-spiece.model",
|
||||
"albert-large-v2": "https://s3.amazonaws.com/models.huggingface.co/bert/albert-large-v2-spiece.model",
|
||||
"albert-xlarge-v2": "https://s3.amazonaws.com/models.huggingface.co/bert/albert-xlarge-v2-spiece.model",
|
||||
|
||||
@@ -0,0 +1,52 @@
|
||||
import unittest
|
||||
|
||||
import torch
|
||||
|
||||
from tests.utils import require_torch, slow
|
||||
from transformers import BartTokenizer, BartModel
|
||||
from transformers.modeling_bart import shift_tokens_right
|
||||
|
||||
DEFAULT_DEVICE = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
|
||||
@require_torch
|
||||
class TestHface(unittest.TestCase):
|
||||
|
||||
@classmethod
|
||||
def setUpClass(cls):
|
||||
source_path = "test.source"
|
||||
cls.lns = [" " + x.rstrip() for x in open(source_path).readlines()][:6]
|
||||
tokenizer = BartTokenizer.from_pretrained('bart-large')
|
||||
dct = tokenizer.batch_encode_plus(cls.lns, max_length=1024, return_tensors="pt", pad_to_max_length=True)
|
||||
cls.ids = dct['input_ids'].to(DEFAULT_DEVICE)
|
||||
cls.enc_mask = dct['attention_mask'].to(DEFAULT_DEVICE)
|
||||
cls.prev_output_tokens = shift_tokens_right(cls.ids, 1).to(DEFAULT_DEVICE)
|
||||
cls.model = BartForConditionalGeneration.from_pretrained('bart-large-cnn').half().to(DEFAULT_DEVICE).half()
|
||||
return cls
|
||||
|
||||
|
||||
|
||||
@classmethod
|
||||
def setUpClass(cls):
|
||||
source_path = "test.source"
|
||||
cls.lns = [" " + x.rstrip() for x in open(source_path).readlines()][:6]
|
||||
tokenizer = BartTokenizer.from_pretrained('bart-large')
|
||||
dct = tokenizer.batch_encode_plus(cls.lns, max_length=100, return_tensors="pt", pad_to_max_length=True)
|
||||
cls.ids = dct['input_ids'].to(DEFAULT_DEVICE)
|
||||
cls.prev_output_tokens = shift_tokens_right(cls.ids, 1).to(DEFAULT_DEVICE)
|
||||
cls.model = BartModel.from_pretrained('bart-large').to(DEFAULT_DEVICE)
|
||||
#cls.lns = pickle_load('/Users/shleifer/transformers_fork/lns.pkl')
|
||||
return cls
|
||||
|
||||
def test_hf_fwd_batch(self):
|
||||
bart = self.model
|
||||
bart.reset_logs()
|
||||
with torch.no_grad():
|
||||
bart(self.ids)
|
||||
try:
|
||||
log_df = bart.combine_logs()
|
||||
#log_df.to_csv('hf_batch_fwd_logs.csv')
|
||||
bart.save_logs('hf_batch_fwd_logs.txt')
|
||||
print(bart.summary)
|
||||
except AttributeError as e:
|
||||
print(e)
|
||||
|
||||
Reference in New Issue
Block a user