Compare commits
10
Commits
v4.1.0
...
mem-prof-bart
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
29a1d7bfe4 | ||
|
|
93d42f7b54 | ||
|
|
2fc31af923 | ||
|
|
e0a500d620 | ||
|
|
8e9fa9bb33 | ||
|
|
e9cc7f2895 | ||
|
|
53f1dcbdb5 | ||
|
|
76e6652391 | ||
|
|
17f3ae3bb8 | ||
|
|
e2931f3860 |
@@ -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",
|
||||
|
||||
File diff suppressed because one or more lines are too long
@@ -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)
|
||||
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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')
|
||||
@@ -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):
|
||||
|
||||
Reference in New Issue
Block a user