Compare commits

...
10 Commits
Author SHA1 Message Date
sshleifer 29a1d7bfe4 boom boom 2020-03-24 00:27:01 -04:00
sshleifer 93d42f7b54 boom boom 2020-03-24 00:26:06 -04:00
sshleifer 2fc31af923 boom boom 2020-03-24 00:23:47 -04:00
sshleifer e0a500d620 boom boom 2020-03-24 00:06:37 -04:00
sshleifer 8e9fa9bb33 boom boom 2020-03-24 00:03:38 -04:00
sshleifer e9cc7f2895 call init 2020-03-23 23:54:44 -04:00
sshleifer 53f1dcbdb5 add example logging statement 2020-03-23 23:47:04 -04:00
sshleifer 76e6652391 add fairseq 2020-03-23 23:22:54 -04:00
sshleifer 17f3ae3bb8 works 2020-03-23 23:13:52 -04:00
sshleifer e2931f3860 boom boom 2020-03-23 22:10:08 -04:00
7 changed files with 297 additions and 2 deletions
+1 -1
View File
@@ -46,7 +46,7 @@ def generate_summaries(lns, out_file, batch_size=8, device=DEFAULT_DEVICE):
def _run_generate():
parser = argparse.ArgumentParser()
parser.add_argument(
"source_path", type=str, help="like cnn_dm/test.source",
"DATA_PATH", type=str, help="like cnn_dm/test.source",
)
parser.add_argument(
"output_path", type=str, help="where to save summaries",
+100
View File
File diff suppressed because one or more lines are too long
+67
View File
@@ -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(
"--DATA_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)
+2
View File
@@ -769,7 +769,9 @@ class BartModel(PretrainedBartModel):
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,
+1
View File
@@ -920,6 +920,7 @@ class PreTrainedModel(nn.Module, ModuleUtilsMixin):
encoder = self.get_encoder()
encoder_outputs = encoder(input_ids, attention_mask=attention_mask)
self.log_mem(f'done encoder, outputs shaped {encoder_outputs[0].shape}')
# create empty decoder_input_ids
input_ids = torch.full(
+122
View File
@@ -0,0 +1,122 @@
import unittest
import torch
import os
from tests.utils import require_torch, slow
from transformers import BartTokenizer, BartModel, BartForConditionalGeneration
from transformers.modeling_bart import shift_tokens_right
from py3nvml.py3nvml import *
DEFAULT_DEVICE = "cuda" if torch.cuda.is_available() else "cpu"
DATA_PATH = "small_test.source"
SAVE_PREFIX = os.getenv('SAVE_PREFIX', '')
from durbango.logging_utils import collect_log_data
from durbango import patch_module_with_memory_mixin
def save_logs_print_mem(bart, save_path):
pth = SAVE_PREFIX + save_path
print(f'*** {pth} ***')
bart.save_logs(pth+'.txt')
bart.save_log_csv(pth+'.csv')
print(bart.summary)
print(f'*** DONE ***')
import py3nvml
class Memtest(unittest.TestCase):
def setUp(self):
if hasattr(self.model, 'reset_logs'): self.model.reset_logs()
self.model.log_mem('start')
torch.cuda.empty_cache()
if torch.cuda.is_available():
nvmlInit()
def tearDown(self) -> None:
try:
nvmlShutdown()
except Exception:
pass
class TestHface(Memtest):
@classmethod
def setUpClass(cls):
cls.lns = [" " + x.rstrip() for x in open(DATA_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').to(DEFAULT_DEVICE)
patch_module_with_memory_mixin(cls.model)
return cls
def test_hf_fwd(self):
bart = self.model
with torch.no_grad():
self.model(self.ids, attention_mask=self.enc_mask, generation_mode=False)
self.model.log_mem('done')
save_logs_print_mem(self.model, 'hf_fwd')
def test_hf_short_generate(self):
self.model.generate(self.ids, attention_mask=self.enc_mask, num_beams=4,
max_length=9, min_length=6,
no_repeat_ngram_size=3,
early_stopping=True,
decoder_start_token_id=2,
)
self.model.log_mem('done')
save_logs_print_mem(self.model, 'hf_short_generate')
@slow
def test_hf_generate(self):
self.model.generate(self.ids, attention_mask=self.enc_mask, num_beams=4, max_length=140, min_length=56,
no_repeat_ngram_size=3,
early_stopping=True,
decoder_start_token_id=2,
)
self.model.log_mem('done')
save_logs_print_mem(self.model, 'hf_generate')
try:
import fairseq
HAS_FAIRSEQ = True
except ImportError:
HAS_FAIRSEQ = False
class TestFairseq(Memtest):
@classmethod
def setUpClass(cls):
import fairseq
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.prev_output_tokens = shift_tokens_right(cls.ids, 1).to(DEFAULT_DEVICE)
cls.model = torch.hub.load('pytorch/fairseq', 'bart.large.cnn').eval().to(DEFAULT_DEVICE)
patch_module_with_memory_mixin(cls.model)
return cls
def test_fs_fwd(self):
bart = self.model
with torch.no_grad():
bart.model(self.ids, None, self.prev_output_tokens)
bart.log_mem('done')
save_logs_print_mem(bart, 'fs_fwd')
def test_fs_short_gen(self):
bart = self.model
bart.sample(self.lns, beam=4, lenpen=2.0, max_len_b=7, min_len=5, no_repeat_ngram_size=3)
bart.log_mem('done')
save_logs_print_mem(bart, 'fs_short_generate')
@slow
def test_fs_gen(self):
bart = self.model
bart.sample(self.lns, beam=4, lenpen=2.0, max_len_b=140, min_len=55, no_repeat_ngram_size=3)
bart.log_mem('done')
save_logs_print_mem(bart, 'fs_generate')
+4 -1
View File
@@ -279,7 +279,10 @@ class BartHeadTests(unittest.TestCase):
bos_token_id=0,
)
lm_model = BartForConditionalGeneration(config).to(torch_device)
lm_model.apply(patch_module_with_memory_mixin)
lm_model.eval()
lm_model.log_mem()
lm_model.model.decoder.log_mem()
max_length = 5
new_input_ids = lm_model.generate(
@@ -392,7 +395,7 @@ def _long_tensor(tok_lst):
TOLERANCE = 1e-4
from durbango.logging_utils import patch_module_with_memory_mixin
@require_torch
class BartModelIntegrationTests(unittest.TestCase):