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_PRETRAINED_CONFIG_ARCHIVE_MAP = {
|
||||||
"albert-base-v1": "https://s3.amazonaws.com/models.huggingface.co/bert/albert-base-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-v1-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-v1-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-v1-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-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-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",
|
"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_PRETRAINED_MODEL_ARCHIVE_MAP = {
|
||||||
"albert-base-v1": "https://s3.amazonaws.com/models.huggingface.co/bert/albert-base-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-v1-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-v1-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-v1-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-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-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",
|
"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 .configuration_bart import BartConfig
|
||||||
from .file_utils import add_start_docstrings, add_start_docstrings_to_callable
|
from .file_utils import add_start_docstrings, add_start_docstrings_to_callable
|
||||||
from .modeling_utils import PreTrainedModel, create_position_ids_from_input_ids
|
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__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
@@ -109,7 +110,7 @@ def _prepare_bart_decoder_inputs(
|
|||||||
return decoder_input_ids, decoder_attn_mask
|
return decoder_input_ids, decoder_attn_mask
|
||||||
|
|
||||||
|
|
||||||
class PretrainedBartModel(PreTrainedModel):
|
class PretrainedBartModel(PreTrainedModel, LoggingMixin):
|
||||||
config_class = BartConfig
|
config_class = BartConfig
|
||||||
base_model_prefix = "model"
|
base_model_prefix = "model"
|
||||||
pretrained_model_archive_map = BART_PRETRAINED_MODEL_ARCHIVE_MAP
|
pretrained_model_archive_map = BART_PRETRAINED_MODEL_ARCHIVE_MAP
|
||||||
@@ -185,9 +186,10 @@ def make_padding_mask(input_ids, padding_idx=1):
|
|||||||
|
|
||||||
|
|
||||||
# Helper Modules
|
# Helper Modules
|
||||||
|
from durbango.torch_utils import get_shapes
|
||||||
|
|
||||||
|
|
||||||
class EncoderLayer(nn.Module):
|
class EncoderLayer(nn.Module, LoggingMixin):
|
||||||
def __init__(self, config: BartConfig):
|
def __init__(self, config: BartConfig):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.embed_dim = config.d_model
|
self.embed_dim = config.d_model
|
||||||
@@ -216,7 +218,8 @@ class EncoderLayer(nn.Module):
|
|||||||
encoded output of shape `(seq_len, batch, embed_dim)`
|
encoded output of shape `(seq_len, batch, embed_dim)`
|
||||||
"""
|
"""
|
||||||
residual = x
|
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 = F.dropout(x, p=self.dropout, training=self.training)
|
||||||
x = residual + x
|
x = residual + x
|
||||||
x = self.self_attn_layer_norm(x)
|
x = self.self_attn_layer_norm(x)
|
||||||
@@ -230,8 +233,8 @@ class EncoderLayer(nn.Module):
|
|||||||
x = self.final_layer_norm(x)
|
x = self.final_layer_norm(x)
|
||||||
return x, attn_weights
|
return x, attn_weights
|
||||||
|
|
||||||
|
import gc
|
||||||
class BartEncoder(nn.Module):
|
class BartEncoder(nn.Module, LoggingMixin):
|
||||||
"""
|
"""
|
||||||
Transformer encoder consisting of *config.encoder_layers* self attention layers. Each layer
|
Transformer encoder consisting of *config.encoder_layers* self attention layers. Each layer
|
||||||
is a :class:`EncoderLayer`.
|
is a :class:`EncoderLayer`.
|
||||||
@@ -282,16 +285,19 @@ class BartEncoder(nn.Module):
|
|||||||
attention_mask = attention_mask.eq(0)
|
attention_mask = attention_mask.eq(0)
|
||||||
|
|
||||||
inputs_embeds = self.embed_tokens(input_ids)
|
inputs_embeds = self.embed_tokens(input_ids)
|
||||||
embed_pos = self.embed_positions(input_ids)
|
x = inputs_embeds + self.embed_positions(input_ids)
|
||||||
x = inputs_embeds + embed_pos
|
|
||||||
x = self.layernorm_embedding(x)
|
x = self.layernorm_embedding(x)
|
||||||
x = F.dropout(x, p=self.dropout, training=self.training)
|
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
|
# B x T x C -> T x B x C
|
||||||
x = x.transpose(0, 1)
|
x = x.transpose(0, 1)
|
||||||
|
self.log_mem('encoder: starting_loop')
|
||||||
encoder_states, all_attentions = [], []
|
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:
|
if self.output_hidden_states:
|
||||||
encoder_states.append(x)
|
encoder_states.append(x)
|
||||||
# add LayerDrop (see https://arxiv.org/abs/1909.11556 for description)
|
# 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
|
if self.training and (dropout_probability < self.layerdrop): # skip the layer
|
||||||
attn = None
|
attn = None
|
||||||
else:
|
else:
|
||||||
x, attn = encoder_layer(x, attention_mask)
|
x, _ = encoder_layer(x, attention_mask)
|
||||||
|
assert len(encoder_states) == 0
|
||||||
if self.output_attentions:
|
assert len(all_attentions) == 0
|
||||||
all_attentions.append(attn)
|
#self.log_mem(f'x: {x.shape}, attn: {attn.shape}')
|
||||||
|
self.log_mem(f'Encoder: called layer {i}')
|
||||||
if self.output_hidden_states:
|
|
||||||
encoder_states.append(x)
|
|
||||||
|
|
||||||
encoder_states = [hidden_state.transpose(0, 1) for hidden_state in encoder_states]
|
encoder_states = [hidden_state.transpose(0, 1) for hidden_state in encoder_states]
|
||||||
return x, encoder_states, all_attentions
|
return x, encoder_states, all_attentions
|
||||||
|
|
||||||
|
|
||||||
class DecoderLayer(nn.Module):
|
class DecoderLayer(nn.Module, LoggingMixin):
|
||||||
def __init__(self, config: BartConfig):
|
def __init__(self, config: BartConfig):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.embed_dim = config.d_model
|
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
|
) # 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
|
Transformer decoder consisting of *config.decoder_layers* layers. Each layer
|
||||||
is a :class:`DecoderLayer`.
|
is a :class:`DecoderLayer`.
|
||||||
@@ -435,6 +438,8 @@ class BartDecoder(nn.Module):
|
|||||||
encoder_padding_mask = encoder_padding_mask.eq(0)
|
encoder_padding_mask = encoder_padding_mask.eq(0)
|
||||||
|
|
||||||
# embed positions
|
# embed positions
|
||||||
|
|
||||||
|
self.log_mem('decoder: embedded positions')
|
||||||
positions = self.embed_positions(input_ids, generation_mode=generation_mode)
|
positions = self.embed_positions(input_ids, generation_mode=generation_mode)
|
||||||
|
|
||||||
if generation_mode:
|
if generation_mode:
|
||||||
@@ -443,6 +448,7 @@ class BartDecoder(nn.Module):
|
|||||||
assert input_ids.ne(self.padding_idx).any()
|
assert input_ids.ne(self.padding_idx).any()
|
||||||
|
|
||||||
x = self.embed_tokens(input_ids)
|
x = self.embed_tokens(input_ids)
|
||||||
|
self.log_mem('decoder: embedded tokens')
|
||||||
x += positions
|
x += positions
|
||||||
|
|
||||||
x = self.layernorm_embedding(x)
|
x = self.layernorm_embedding(x)
|
||||||
@@ -464,6 +470,7 @@ class BartDecoder(nn.Module):
|
|||||||
x, layer_self_attn, layer_past = decoder_layer(
|
x, layer_self_attn, layer_past = decoder_layer(
|
||||||
x, encoder_hidden_states, encoder_padding_mask, layer_state=layer_state, attention_mask=combined_mask,
|
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:
|
if self.output_past:
|
||||||
next_decoder_cache.append(layer_past.copy())
|
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)
|
return x, next_cache, all_hidden_states, list(all_self_attns)
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
def _reorder_buffer(attn_cache, new_order):
|
def _reorder_buffer(attn_cache, new_order):
|
||||||
for k, input_buffer_k in attn_cache.items():
|
for k, input_buffer_k in attn_cache.items():
|
||||||
if input_buffer_k is not None:
|
if input_buffer_k is not None:
|
||||||
@@ -490,8 +498,8 @@ def _reorder_buffer(attn_cache, new_order):
|
|||||||
return attn_cache
|
return attn_cache
|
||||||
|
|
||||||
|
|
||||||
class SelfAttention(nn.Module):
|
class SelfAttention(nn.Module, LoggingMixin):
|
||||||
"""Multi-headed attention from 'Attention Is All You Need' paper"""
|
"""Multi-headed attention from "Attention Is All You Need"""
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
@@ -519,11 +527,16 @@ class SelfAttention(nn.Module):
|
|||||||
def _shape(self, tensor, dim_0, bsz):
|
def _shape(self, tensor, dim_0, bsz):
|
||||||
return tensor.contiguous().view(dim_0, bsz * self.num_heads, self.head_dim).transpose(0, 1)
|
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(
|
def forward(
|
||||||
self,
|
self,
|
||||||
query,
|
query,
|
||||||
key: Optional[Tensor],
|
key: Optional[Tensor],
|
||||||
key_padding_mask: Optional[Tensor] = None,
|
key_padding_mask: Optional[Tensor] = None,
|
||||||
|
update_layer_state=True,
|
||||||
layer_state: Optional[Dict[str, Optional[Tensor]]] = None,
|
layer_state: Optional[Dict[str, Optional[Tensor]]] = None,
|
||||||
attn_mask: Optional[Tensor] = None,
|
attn_mask: Optional[Tensor] = None,
|
||||||
) -> Tuple[Tensor, Optional[Tensor]]:
|
) -> Tuple[Tensor, Optional[Tensor]]:
|
||||||
@@ -544,6 +557,7 @@ class SelfAttention(nn.Module):
|
|||||||
layer_state = {}
|
layer_state = {}
|
||||||
|
|
||||||
q = self.q_proj(query) * self.scaling
|
q = self.q_proj(query) * self.scaling
|
||||||
|
self.log_mem('\tq_proj')
|
||||||
if static_kv:
|
if static_kv:
|
||||||
if key is None:
|
if key is None:
|
||||||
k = v = None
|
k = v = None
|
||||||
@@ -554,29 +568,39 @@ class SelfAttention(nn.Module):
|
|||||||
k = self.k_proj(query)
|
k = self.k_proj(query)
|
||||||
v = self.v_proj(query)
|
v = self.v_proj(query)
|
||||||
|
|
||||||
|
|
||||||
q = self._shape(q, tgt_len, bsz)
|
q = self._shape(q, tgt_len, bsz)
|
||||||
|
self.log_mem(f'\tq_reshape -> {q.shape}')
|
||||||
if k is not None:
|
if k is not None:
|
||||||
k = self._shape(k, -1, bsz)
|
k = self._shape(k, -1, bsz)
|
||||||
|
self.log_mem(f'\t done reshaping k,v ->, {k.shape}')
|
||||||
if v is not None:
|
if v is not None:
|
||||||
v = self._shape(v, -1, bsz)
|
v = self._shape(v, -1, bsz)
|
||||||
|
|
||||||
|
|
||||||
if saved_state is not None:
|
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)
|
k, v, key_padding_mask = self._use_saved_state(k, v, saved_state, key_padding_mask, static_kv, bsz)
|
||||||
|
|
||||||
# Update cache
|
# Update cache
|
||||||
layer_state[self.cache_key] = {
|
if update_layer_state:
|
||||||
"prev_key": k.view(bsz, self.num_heads, -1, self.head_dim),
|
layer_state[self.cache_key] = {
|
||||||
"prev_value": v.view(bsz, self.num_heads, -1, self.head_dim),
|
"prev_key": k.view(bsz, self.num_heads, -1, self.head_dim),
|
||||||
"prev_key_padding_mask": key_padding_mask if not static_kv else None,
|
"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
|
assert k is not None
|
||||||
src_len = k.size(1)
|
src_len = k.size(1)
|
||||||
|
self.log_mem('\t attn: before BMM(q,k)')
|
||||||
attn_weights = torch.bmm(q, k.transpose(1, 2))
|
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)
|
assert attn_weights.size() == (bsz * self.num_heads, tgt_len, src_len)
|
||||||
|
|
||||||
if attn_mask is not None:
|
if attn_mask is not None:
|
||||||
attn_weights = attn_weights.view(bsz, self.num_heads, tgt_len, src_len) + attn_mask
|
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)
|
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.
|
# 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
|
key_padding_mask = None
|
||||||
assert key_padding_mask is None or key_padding_mask.size()[:2] == (bsz, src_len,)
|
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)
|
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(key_padding_mask.unsqueeze(1).unsqueeze(2), float("-inf"))
|
||||||
attn_weights = attn_weights.masked_fill(reshaped, 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 = attn_weights.view(bsz * self.num_heads, tgt_len, src_len)
|
||||||
attn_weights = F.softmax(attn_weights, dim=-1)
|
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,)
|
attn_probs = F.dropout(attn_weights, p=self.dropout, training=self.training,)
|
||||||
|
|
||||||
|
|
||||||
assert v is not None
|
assert v is not None
|
||||||
attn_output = torch.bmm(attn_probs, v)
|
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)
|
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)
|
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_output = self.out_proj(attn_output)
|
||||||
attn_weights = attn_weights.view(bsz, self.num_heads, tgt_len, src_len)
|
self.log_mem('\t attn: done out_proj')
|
||||||
return attn_output, attn_weights
|
#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):
|
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)
|
# 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
|
return new_key_padding_mask
|
||||||
|
|
||||||
|
|
||||||
class BartClassificationHead(nn.Module):
|
class BartClassificationHead(nn.Module, LoggingMixin):
|
||||||
"""Head for sentence-level classification tasks."""
|
"""Head for sentence-level classification tasks."""
|
||||||
|
|
||||||
# This can trivially be shared with RobertaClassificationHead
|
# This can trivially be shared with RobertaClassificationHead
|
||||||
@@ -727,6 +756,8 @@ def _filter_out_falsey_values(tup) -> Tuple:
|
|||||||
|
|
||||||
# Public API
|
# Public API
|
||||||
|
|
||||||
|
import time
|
||||||
|
import pandas as pd
|
||||||
|
|
||||||
@add_start_docstrings(
|
@add_start_docstrings(
|
||||||
"The bare BART Model outputting raw hidden-states without any specific head on top.", BART_START_DOCSTRING,
|
"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.encoder = BartEncoder(config, self.shared)
|
||||||
self.decoder = BartDecoder(config, self.shared)
|
self.decoder = BartDecoder(config, self.shared)
|
||||||
|
|
||||||
|
|
||||||
self.init_weights()
|
self.init_weights()
|
||||||
|
|
||||||
@add_start_docstrings_to_callable(BART_INPUTS_DOCSTRING)
|
@add_start_docstrings_to_callable(BART_INPUTS_DOCSTRING)
|
||||||
@@ -758,6 +790,9 @@ class BartModel(PretrainedBartModel):
|
|||||||
):
|
):
|
||||||
|
|
||||||
# make masks if user doesn't supply
|
# 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:
|
if not generation_mode:
|
||||||
decoder_input_ids, decoder_attention_mask = _prepare_bart_decoder_inputs(
|
decoder_input_ids, decoder_attention_mask = _prepare_bart_decoder_inputs(
|
||||||
self.config,
|
self.config,
|
||||||
@@ -767,9 +802,6 @@ class BartModel(PretrainedBartModel):
|
|||||||
mask_dtype=self.shared.weight.dtype,
|
mask_dtype=self.shared.weight.dtype,
|
||||||
)
|
)
|
||||||
assert decoder_input_ids is not None
|
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 consists of (dec_features, layer_state, dec_hidden, dec_attn)
|
||||||
decoder_outputs = self.decoder(
|
decoder_outputs = self.decoder(
|
||||||
decoder_input_ids,
|
decoder_input_ids,
|
||||||
@@ -804,10 +836,10 @@ class BartForConditionalGeneration(PretrainedBartModel):
|
|||||||
|
|
||||||
def __init__(self, config: BartConfig):
|
def __init__(self, config: BartConfig):
|
||||||
super().__init__(config)
|
super().__init__(config)
|
||||||
# if base_model is None:
|
# if base_model is Nones:
|
||||||
base_model = BartModel(config)
|
#self.log_mem('pre-init')
|
||||||
self.model = base_model
|
self.model = BartModel(config)
|
||||||
self.lm_head = _make_linear_from_emb(self.model.shared)
|
#self.lm_head = _make_linear_from_emb(self.model.shared)
|
||||||
|
|
||||||
def tie_weights(self):
|
def tie_weights(self):
|
||||||
pass # hack to prevent changing lm_head.out_features. The input and output embeddings are still the same.
|
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()
|
tokenizer.decode(predictions).split()
|
||||||
# ['good', 'great', 'all', 'really', 'very']
|
# ['good', 'great', 'all', 'really', 'very']
|
||||||
"""
|
"""
|
||||||
|
self.model.log_mem('before BartModel.forward')
|
||||||
outputs = self.model(
|
outputs = self.model(
|
||||||
input_ids,
|
input_ids,
|
||||||
attention_mask=attention_mask,
|
attention_mask=attention_mask,
|
||||||
@@ -875,7 +908,10 @@ class BartForConditionalGeneration(PretrainedBartModel):
|
|||||||
decoder_cached_states=decoder_cached_states,
|
decoder_cached_states=decoder_cached_states,
|
||||||
generation_mode=generation_mode,
|
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
|
outputs = (lm_logits,) + outputs[1:] # Add hidden states and attention if they are here
|
||||||
if lm_labels is not None:
|
if lm_labels is not None:
|
||||||
loss_fct = nn.CrossEntropyLoss()
|
loss_fct = nn.CrossEntropyLoss()
|
||||||
@@ -885,6 +921,7 @@ class BartForConditionalGeneration(PretrainedBartModel):
|
|||||||
|
|
||||||
return outputs
|
return outputs
|
||||||
|
|
||||||
|
|
||||||
def prepare_inputs_for_generation(self, decoder_input_ids, past, attention_mask, **kwargs):
|
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"
|
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
|
encoder_outputs, decoder_cached_states = past, None
|
||||||
else:
|
else:
|
||||||
encoder_outputs, decoder_cached_states = past
|
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 {
|
return {
|
||||||
"input_ids": None, # encoder_outputs is defined. input_ids not needed
|
"input_ids": None, # encoder_outputs is defined. input_ids not needed
|
||||||
"encoder_outputs": encoder_outputs,
|
"encoder_outputs": encoder_outputs,
|
||||||
@@ -932,7 +970,7 @@ class BartForConditionalGeneration(PretrainedBartModel):
|
|||||||
return self.model.encoder
|
return self.model.encoder
|
||||||
|
|
||||||
def get_output_embeddings(self):
|
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(
|
@add_start_docstrings(
|
||||||
|
|||||||
@@ -814,9 +814,10 @@ class T5Model(T5PreTrainedModel):
|
|||||||
|
|
||||||
return decoder_outputs + encoder_outputs
|
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)
|
@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"""
|
r"""
|
||||||
**lm_labels**: (`optional`) ``torch.LongTensor`` of shape ``(batch_size, sequence_length)``:
|
**lm_labels**: (`optional`) ``torch.LongTensor`` of shape ``(batch_size, sequence_length)``:
|
||||||
Labels for computing the masked language modeling loss.
|
Labels for computing the masked language modeling loss.
|
||||||
|
|||||||
@@ -896,6 +896,17 @@ class PreTrainedModel(nn.Module, ModuleUtilsMixin):
|
|||||||
effective_batch_size = batch_size
|
effective_batch_size = batch_size
|
||||||
effective_batch_mult = 1
|
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
|
# Expand input ids if num_beams > 1 or num_return_sequences > 1
|
||||||
if num_return_sequences > 1 or num_beams > 1:
|
if num_return_sequences > 1 or num_beams > 1:
|
||||||
input_ids_len = input_ids.shape[-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
|
effective_batch_size * num_beams, input_ids_len
|
||||||
) # shape: (batch_size * num_return_sequences * num_beams, cur_len)
|
) # shape: (batch_size * num_return_sequences * num_beams, cur_len)
|
||||||
|
|
||||||
|
|
||||||
if self.config.is_encoder_decoder:
|
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
|
# create empty decoder_input_ids
|
||||||
input_ids = torch.full(
|
input_ids = torch.full(
|
||||||
(effective_batch_size * num_beams, 1),
|
(effective_batch_size * num_beams, 1),
|
||||||
@@ -929,10 +932,16 @@ class PreTrainedModel(nn.Module, ModuleUtilsMixin):
|
|||||||
device=next(self.parameters()).device,
|
device=next(self.parameters()).device,
|
||||||
)
|
)
|
||||||
cur_len = 1
|
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:
|
else:
|
||||||
encoder_outputs = None
|
encoder_outputs = None
|
||||||
cur_len = input_ids.shape[-1]
|
cur_len = input_ids.shape[-1]
|
||||||
|
|
||||||
|
|
||||||
if num_beams > 1:
|
if num_beams > 1:
|
||||||
output = self._generate_beam_search(
|
output = self._generate_beam_search(
|
||||||
input_ids,
|
input_ids,
|
||||||
|
|||||||
@@ -29,10 +29,10 @@ VOCAB_FILES_NAMES = {"vocab_file": "spiece.model"}
|
|||||||
|
|
||||||
PRETRAINED_VOCAB_FILES_MAP = {
|
PRETRAINED_VOCAB_FILES_MAP = {
|
||||||
"vocab_file": {
|
"vocab_file": {
|
||||||
"albert-base-v1": "https://s3.amazonaws.com/models.huggingface.co/bert/albert-base-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-v1-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-v1-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-v1-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-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-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",
|
"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