delete run_distiller -> main

This commit is contained in:
sshleifer
2020-05-21 01:20:11 -04:00
parent 2c2a10e5d7
commit 21a2cc0eb1
2 changed files with 28 additions and 32 deletions
+18 -26
View File
@@ -78,6 +78,13 @@ class SummarizationTrainer(BaseTransformer):
"test": self.hparams.n_test,
}
self.n_obs = {k: v if v >= 0 else None for k, v in base_nobs.items()}
self.freeze_stuff()
def freeze_stuff(self):
freeze_part(self.model.model.shared)
for d in [self.model.model.encoder, self.model.model.decoder]:
freeze_part(d.embed_positions)
freeze_part(d.embed_tokens)
@property
def metrics_df(self):
@@ -240,6 +247,7 @@ def get_layers_to_copy(n_to_get, tot):
class SummarizationDistiller(SummarizationTrainer):
loss_names = ["loss", "ce_loss", "mlm_loss"]
def __init__(self, hparams):
# Dump empty student model at a path, then call from_pretrained on it
@@ -262,29 +270,21 @@ class SummarizationDistiller(SummarizationTrainer):
copy_layers(teacher.model.decoder.layers, student.model.decoder.layers, d_layers_to_copy)
copy_layers(teacher.model.encoder.layers, student.model.encoder.layers, e_layers_to_copy)
if e_layers_to_copy == 12:
freeze_part(self.model.model.encoder)
teacher.model.encoder = None
# Path(hparams.model_name_or_path).mkdir(exist_ok=True)
tokenizer = BartTokenizer.from_pretrained("bart-large")
super().__init__(hparams, model=student, config=student_cfg, tokenizer=tokenizer)
assert len(self.model.model.decoder.layers) == len(d_layers_to_copy)
self.model.teacher = teacher
# self.teacher = teacher
freeze_part(self.model.teacher)
self.ce_loss_fct = nn.KLDivLoss(reduction="batchmean")
self.freeze_stuff()
self.temperature = 2.0
self.alpha_mlm = hparams.alpha_mlm
self.alpha_ce = hparams.alpha_ce
def freeze_stuff(self):
freeze_part(self.model.model.encoder)
freeze_part(self.model.model.shared)
d = self.model.model.decoder
freeze_part(d.embed_positions)
freeze_part(d.embed_tokens)
self.model.teacher.encoder = None
freeze_part(self.model.teacher)
def _step(self, batch):
# assert is_frozen(self.model.teacher)
pad_token_id = self.tokenizer.pad_token_id
@@ -368,27 +368,19 @@ class SummarizationDistiller(SummarizationTrainer):
parser.add_argument(
"--student_encoder_layers", default=12, type=int, required=False,
)
parser.add_argument(
"--distilled_ds", action="store_true", default=False,
)
return parser
def main(args):
# If output_dir not provided, a folder will be generated in pwd
if not args.output_dir:
args.output_dir = os.path.join("./results", f"{args.task}_{time.strftime('%Y%m%d_%H%M%S')}",)
os.makedirs(args.output_dir)
model: pl.LightningModule = SummarizationTrainer(args)
trainer: pl.Trainer = generic_train(model, args)
checkpoints = list(sorted(glob.glob(os.path.join(args.output_dir, "checkpointepoch=*.ckpt"), recursive=True)))
model = model.load_from_checkpoint(checkpoints[-1])
trainer.test(model)
def run_distiller(args):
if not args.output_dir:
args.output_dir = os.path.join("./results", f"{args.task}_{time.strftime('%Y%m%d_%H%M%S')}",)
os.makedirs(args.output_dir)
model: pl.LightningModule = SummarizationDistiller(args)
module_cls = SummarizationTrainer if args.distilled_ds else SummarizationDistiller
model: SummarizationTrainer = module_cls(args)
trainer: pl.Trainer = generic_train(model, args, early_stopping_callback=True)
checkpoints = list(sorted(glob.glob(os.path.join(args.output_dir, "checkpointepoch=*.ckpt"), recursive=True)))
model = model.load_from_checkpoint(checkpoints[-1])
@@ -402,4 +394,4 @@ if __name__ == "__main__":
parser = SummarizationDistiller.add_model_specific_args(parser, os.getcwd())
args = parser.parse_args()
run_distiller(args)
main(args)
@@ -14,7 +14,7 @@ from durbango import DEFAULT_DEVICE, pickle_load
from transformers import BartTokenizer
from .evaluate_cnn import run_generate
from .finetune import main, run_distiller
from .finetune import main
from .utils import SummarizationDataset, summaries_for_file
@@ -29,6 +29,7 @@ CHEAP_ARGS = {
"val_check_interval": 1.0,
"output_dir": "",
"fp16": False,
"distilled_ds": False,
"fp16_opt_level": "O1",
"n_gpu": 0,
"n_tpu_cores": 0,
@@ -102,8 +103,12 @@ class TestBartExamples(unittest.TestCase):
tmp_dir = make_test_data_dir()
output_dir = tempfile.mkdtemp(prefix="output_")
args_d.update(
data_dir=tmp_dir, model_type="bart", train_batch_size=2, eval_batch_size=2,
data_dir=tmp_dir,
model_type="bart",
train_batch_size=2,
eval_batch_size=2,
num_train_epochs=2,
distilled_ds=True,
output_dir=output_dir,
)
main(argparse.Namespace(**args_d))
@@ -120,7 +125,7 @@ class TestBartExamples(unittest.TestCase):
"test_generations_.txt",
"test_generations_1.txt",
"val_generations_0.txt",
#"val_generations_1.txt",
# "val_generations_1.txt",
"val_generations_.txt",
# "test"
}
@@ -137,7 +142,6 @@ class TestBartExamples(unittest.TestCase):
train_batch_size=1,
eval_batch_size=2,
num_train_epochs=2,
n_gpu=0,
output_dir=output_dir,
do_predict=True,
@@ -147,7 +151,7 @@ class TestBartExamples(unittest.TestCase):
alpha_ce=0.8,
val_check_interval=0.5,
)
run_distiller(argparse.Namespace(**args_d))
main(argparse.Namespace(**args_d))
contents = os.listdir(output_dir)
expected_contents = {
"checkpointepoch=0.ckpt",
@@ -189,7 +193,7 @@ class TestBartExamples(unittest.TestCase):
model_name_or_path="student",
teacher="bart-large-cnn",
)
run_distiller(argparse.Namespace(**args_d))
main(argparse.Namespace(**args_d))
contents = os.listdir(output_dir)
expected_contents = {
"checkpointepoch=0.ckpt",