Compare commits

...
Author SHA1 Message Date
sshleifer a07fcff087 boom boom 2020-06-25 02:10:13 +00:00
sshleifer 243f2d98f4 boom boom 2020-06-24 20:44:27 +00:00
sshleifer 7c6802d216 boom boom 2020-06-23 16:10:00 -04:00
sshleifer 6dd3122ad4 tests pass 2020-06-23 14:35:54 -04:00
sshleifer 3a8f32d919 pl 0.8.1 compat 2020-06-19 14:46:51 +00:00
sshleifer 29d1eb57ac Support top, bottom init 2020-06-19 09:43:20 -04:00
sshleifer 02418a591b boom boom 2020-06-18 21:38:13 -04:00
sshleifer be95182ee1 half-merge 2020-06-18 21:37:14 -04:00
sshleifer 970fe3d833 smaller distilbart-clean delta 2020-06-17 16:03:49 -04:00
sshleifer 97e29a9cf3 allstats-option 2020-06-15 21:38:08 -04:00
sshleifer 785a1b6076 RougeTracker 2020-06-15 21:13:36 -04:00
sshleifer 9ef5581d5e boom boom 2020-06-15 18:05:27 -04:00
sshleifer 657b7fdada boom boom 2020-06-15 17:57:13 -04:00
sshleifer 3a5999f5d4 boom boom 2020-06-15 17:50:28 -04:00
sshleifer 0efb2830f4 turn off ES callback 2020-06-15 17:43:01 -04:00
sshleifer b5099747f9 boom boom 2020-06-15 17:36:56 -04:00
sshleifer 34fd62a5ea init with copy 2020-06-15 17:36:12 -04:00
sshleifer e167e9fc74 boom boom 2020-06-15 17:19:30 -04:00
sshleifer 1df5e7af4a odir fix 2020-06-15 17:10:53 -04:00
sshleifer 544f543861 rouge_tracker_csvs 2020-06-15 16:56:58 -04:00
sshleifer 10970ec650 fixed dec_mask, maybe 2020-06-15 12:42:54 -04:00
sshleifer 78b1c850f2 passing, unsolved 2020-06-15 12:24:39 -04:00
sshleifer ac8d301786 boom boom 2020-06-14 21:56:11 -04:00
sshleifer 9f32be4fbc boom boom 2020-06-14 21:54:18 -04:00
sshleifer 74c8ac0f84 fixes 2020-06-14 21:51:09 -04:00
sshleifer 1d5e77d87b test_passing 2020-06-14 21:41:49 -04:00
sshleifer 47b3638593 boom boom 2020-06-14 21:18:15 -04:00
sshleifer f7c95805cf boom boom 2020-06-14 21:01:52 -04:00
sshleifer 0169cbbe49 all but t5 passing 2020-06-14 21:01:41 -04:00
sshleifer 8e4a74b1cb borked 2020-06-14 20:46:20 -04:00
sshleifer 86034ae640 Merge branch 'distilbart-clean' into theseus 2020-06-14 19:51:25 -04:00
sshleifer 5283878c9f Move stuff to utils 2020-06-14 19:17:31 -04:00
sshleifer 6bdfb14664 boom boom 2020-06-14 19:07:15 -04:00
sshleifer 99de2c3c3c pass through logger 2020-06-14 19:01:15 -04:00
sshleifer 9e95429843 docs 2020-06-14 18:20:02 -04:00
sshleifer c9597fef42 docs 2020-06-14 18:15:28 -04:00
sshleifer c34b886cd1 Wandb logger 2020-06-14 18:11:22 -04:00
sshleifer f179e7b37c Allow wandb logger 2020-06-14 17:58:41 -04:00
sshleifer 969c271bd7 more honest docs 2020-06-14 20:18:43 +00:00
sshleifer 9f98a6ac15 docs 2020-06-14 16:07:25 -04:00
sshleifer 4606bda735 Bash cleanup 2020-06-14 19:47:04 +00:00
sshleifer eb84b9d5ee boom boom 2020-06-14 15:36:29 -04:00
sshleifer 5526be3499 boom boom 2020-06-14 15:35:47 -04:00
sshleifer 6bf996f58d boom boom 2020-06-14 15:31:46 -04:00
sshleifer 35a82ee87f boom boom 2020-06-14 15:01:50 -04:00
sshleifer 379e8c7dba Cleanup 2020-06-14 15:00:35 -04:00
sshleifer 3e62d96b14 add git-python requirement 2020-06-14 14:28:20 -04:00
sshleifer 350eeb7cd2 style 2020-06-14 13:17:09 -04:00
sshleifer c06e38add3 boom boom 2020-06-14 13:15:17 -04:00
sshleifer d94b09d285 started 2020-06-13 20:27:10 -04:00
sshleifer eab9779732 better mask logic 2020-06-13 20:12:11 -04:00
sshleifer 4a6b23ee40 boom boom 2020-06-13 20:11:41 -04:00
sshleifer e2d45445c3 style 2020-06-13 19:55:14 -04:00
sshleifer de7035d9f6 style 2020-06-13 19:49:35 -04:00
sshleifer abee3e2ecd Merge branch 'master' into distilbart-clean 2020-06-13 16:40:30 -04:00
sshleifer 0274338aff original tests pass 2020-06-13 16:40:04 -04:00
sshleifer e67ad41ae7 boom boom 2020-06-13 15:33:00 -04:00
sshleifer c5a1de5b75 boom boom 2020-06-13 15:30:03 -04:00
sshleifer af3f86fd6a remove some cruft 2020-06-11 23:58:34 -04:00
sshleifer fb6eb096ec Failing but minimal 2020-06-11 23:51:55 -04:00
sshleifer bdf0b1ab88 boom boom 2020-06-11 23:45:32 -04:00
sshleifer 9d99761b1b passing 2020-06-11 22:55:51 -04:00
sshleifer 3f3040ae12 boom boom 2020-06-11 10:51:05 -04:00
sshleifer 4224ac999d boom boom 2020-06-11 09:29:05 -04:00
sshleifer c7d05f70db boom boom 2020-06-11 09:26:30 -04:00
sshleifer ac47dbdfdf boom boom 2020-06-11 09:24:46 -04:00
sshleifer 10b092ea27 boom boom 2020-06-11 09:22:15 -04:00
sshleifer e5728e4ddb boom boom 2020-06-11 02:31:05 +00:00
sshleifer 5a281449df fp16_ever=False 2020-06-11 02:29:48 +00:00
sshleifer 585bb576ba passing cpu 2020-06-10 19:47:02 -04:00
sshleifer 6aa28bd1f5 ignore student 2020-06-10 23:27:33 +00:00
sshleifer a2f2d5da31 undo some chg 2020-06-10 23:23:20 +00:00
sshleifer b2b222d0db boom boom 2020-06-10 17:20:26 -04:00
sshleifer 574ffbb1eb boom boom 2020-06-10 17:17:59 -04:00
sshleifer a549aaac3e boom boom 2020-06-10 17:17:03 -04:00
sshleifer bad1b75a91 boom boom 2020-06-10 17:04:13 -04:00
sshleifer 6624b63c24 boom boom 2020-06-10 16:56:42 -04:00
sshleifer 7f3f9fae22 boom boom 2020-06-10 16:36:06 -04:00
sshleifer c553e1214c t5 works 2020-06-09 23:54:39 -04:00
sshleifer d26f5dac45 stuck on t5 mask 2020-06-09 23:35:59 -04:00
sshleifer 913bb8f770 boom boom 2020-06-09 22:53:11 -04:00
sshleifer af05c0e004 boom boom 2020-06-09 22:52:35 -04:00
sshleifer dc785f0d10 boom boom 2020-06-09 22:51:59 -04:00
sshleifer d3b8694cd8 boom boom 2020-06-09 22:50:44 -04:00
sshleifer b5adf48404 boom boom 2020-06-09 22:49:23 -04:00
sshleifer 315ab6be60 boom boom 2020-06-09 22:48:38 -04:00
sshleifer 76377c9854 boom boom 2020-06-09 22:45:35 -04:00
sshleifer a387c057db boom boom 2020-06-09 22:45:03 -04:00
sshleifer 0098872c04 Deleted SummarizationDistiller 2020-06-09 22:41:36 -04:00
sshleifer a79f5fb77e boom boom 2020-06-09 22:09:57 -04:00
sshleifer 0dc32e6f27 boom boom 2020-06-09 21:11:07 -04:00
sshleifer fa7f8185e8 t5 failing 2020-06-09 17:37:38 -04:00
sshleifer 31b4c8cc08 boom boom 2020-06-09 15:31:31 -04:00
sshleifer 318a585f2b boom boom 2020-06-09 17:44:45 +00:00
sshleifer 630af0c3fa boom boom 2020-06-09 13:27:31 -04:00
sshleifer 3c384445a7 boom boom 2020-06-08 21:58:24 +00:00
sshleifer 9f6596ebc9 boom boom 2020-06-08 14:30:54 -04:00
sshleifer 582efc0319 boom boom 2020-06-08 14:00:07 -04:00
sshleifer 3467f1c7f6 boom boom 2020-06-08 17:58:33 +00:00
sshleifer bcb4624d42 boom boom 2020-06-08 13:57:57 -04:00
sshleifer e46e10150d boom boom 2020-06-08 13:29:51 -04:00
sshleifer ecff6f92fa V8 and generate_summaries fixes 2020-06-08 13:14:46 -04:00
sshleifer c59be9b2e0 boom boom 2020-06-08 13:07:01 -04:00
sshleifer b32d283d26 boom boom 2020-06-08 12:17:51 -04:00
sshleifer 7ccd2366de Evaluate_checkpoint 2020-06-08 11:32:29 -04:00
sshleifer c294e5e331 boom boom 2020-06-07 14:23:46 -04:00
sshleifer 3b304afdf7 boom boom 2020-06-07 14:21:49 -04:00
sshleifer 09f199b790 Merge branch 'distilbart' of github.com:sshleifer/transformers_fork into distilbart 2020-06-07 14:15:28 -04:00
sshleifer bb1bffc89e Save top1 2020-06-07 14:15:23 -04:00
sshleifer 47c226ea3f Merge branch 'distilbart' of github.com:sshleifer/transformers_fork into distilbart 2020-06-07 17:48:10 +00:00
sshleifer 3abbf78a48 boom boom 2020-06-07 17:48:06 +00:00
sshleifer 9fc692d9f3 boom boom 2020-06-06 13:40:07 -04:00
sshleifer e146713686 passing 2020-06-06 13:19:46 -04:00
sshleifer 014c061f3d boom boom 2020-06-06 11:19:45 -04:00
sshleifer e446fd24ef boom boom 2020-06-06 13:13:55 +00:00
sshleifer 055167ba12 boom boom 2020-06-06 01:00:50 -04:00
sshleifer 8f68b91e81 boom boom 2020-06-05 20:10:55 -04:00
sshleifer 8b98c4cfa8 boom boom 2020-06-05 20:02:23 -04:00
sshleifer e3a00a402b boom boom 2020-06-05 19:33:35 -04:00
sshleifer 24f483cf06 boom boom 2020-06-05 19:33:12 -04:00
sshleifer 8d4f075245 boom boom 2020-06-05 19:26:14 -04:00
sshleifer 0b98750ecd boom boom 2020-06-05 17:26:11 -04:00
sshleifer cc83d34a7d boom boom 2020-06-05 15:13:24 -04:00
sshleifer a65ab52d74 boom boom 2020-06-05 19:10:02 +00:00
sshleifer 3e47e61a1a Disable cos loss 2020-06-05 12:51:13 -04:00
sshleifer 20aedfbcf7 boom boom 2020-06-04 09:45:47 -04:00
sshleifer 440744a111 Distiller brewer does not break 2020-06-04 09:38:02 -04:00
sshleifer a2ba031b92 boom boom 2020-06-04 07:52:34 -04:00
sshleifer 8f614af5b7 broken 2020-06-04 07:15:34 -04:00
sshleifer df0b28a581 boom boom 2020-06-04 02:09:27 -04:00
sshleifer 8c6769ee41 boom boom 2020-06-03 13:04:02 +00:00
sshleifer b11fc6977e boom boom 2020-06-03 13:02:37 +00:00
sshleifer fb6adb2a47 backup rouge dfs 2020-06-03 12:53:05 +00:00
sshleifer 789f06e9d8 less annoying requirements 2020-06-02 18:59:11 -04:00
sshleifer e17df7499c merged master 2020-06-02 13:45:10 -04:00
sshleifer 82a6b8345f Merge branch 'master' into distilbart 2020-06-02 13:44:40 -04:00
sshleifer 60a90143ee rouge tracker fix 2020-06-02 11:17:09 -04:00
sshleifer 7b9557fbf5 boom boom 2020-05-31 09:16:20 -04:00
sshleifer 105d83d595 undo rename 2020-05-31 08:52:04 -04:00
sshleifer bd361eae11 rename generic_train -> build_trainer 2020-05-31 08:38:36 -04:00
sshleifer 8fb24ce3ce boom boom 2020-05-27 15:49:36 -04:00
sshleifer 6689b594fc boom boom 2020-05-27 15:46:50 -04:00
sshleifer 25edee17ed boom boom 2020-05-27 14:22:36 -04:00
sshleifer 7c52e73dfe boom boom 2020-05-27 14:19:08 -04:00
sshleifer d53236c840 boom boom 2020-05-27 14:17:34 -04:00
sshleifer eb4b3cce7a boom boom 2020-05-27 13:50:07 -04:00
sshleifer 324d4dcfce No wandb if ddp 2020-05-27 12:49:11 -04:00
sshleifer 3f6463a895 attempt logger.name fix 2020-05-27 11:51:49 -04:00
sshleifer 3c34cae08a distrib hacks 2020-05-27 15:47:45 +00:00
sshleifer 175efd864a getstate 2020-05-27 14:33:15 +00:00
sshleifer 516bbcabd0 boom boom 2020-05-27 10:28:34 -04:00
sshleifer 0c3d6d8cbf boom boom 2020-05-27 10:19:14 -04:00
sshleifer 84ea9fac23 Fixed tests 2020-05-27 09:37:52 -04:00
sshleifer 44743e8dff fix ckpt bug maybe 2020-05-27 08:49:18 -04:00
sshleifer 7f1af24a05 freeze decoder 2020-05-27 01:10:20 -04:00
sshleifer 99e7161eb2 boom boom 2020-05-26 21:23:47 -04:00
sshleifer d92a69f223 boom boom 2020-05-26 21:22:08 -04:00
sshleifer 4f19747960 boom boom 2020-05-26 21:20:53 -04:00
sshleifer 9b54356b0b Logging adjustments 2020-05-26 20:08:41 -04:00
sshleifer f750d801ff boom boom 2020-05-26 15:25:28 -04:00
sshleifer 298ec5361c --freeze_encoder 2020-05-26 15:21:00 -04:00
sshleifer ee18f81852 Merge branch 'distilbart' of github.com:sshleifer/transformers_fork into distilbart 2020-05-26 15:05:36 -04:00
sshleifer e7d9dc3d22 dont default teacher 2020-05-26 19:05:30 +00:00
sshleifer d2cd7b4d9d Somehow passing with better freezing logic 2020-05-26 15:01:56 -04:00
sshleifer d3611c9844 Run distiller: no data defaults 2020-05-26 12:00:12 -04:00
sshleifer c10bb9fb23 boom boom 2020-05-25 23:15:20 -04:00
sshleifer 81e939fb54 boom boom 2020-05-25 21:23:39 -04:00
sshleifer 0c3d22ea86 boom boom 2020-05-25 21:20:08 -04:00
sshleifer 63521f19a9 boom boom 2020-05-25 21:18:18 -04:00
sshleifer 56aedacc27 boom boom 2020-05-25 21:15:00 -04:00
sshleifer 0933e118e0 boom boom 2020-05-25 21:07:02 -04:00
sshleifer 12f2d9e6bc early quitting for encoder only loss 2020-05-25 21:01:47 -04:00
sshleifer 20a59685a9 passing with enc mse 2020-05-25 20:01:04 -04:00
sshleifer f5e697bbe6 Dont require model_name_or_path 2020-05-25 19:31:17 -04:00
sshleifer db52f3708f save fewer checkpoints 2020-05-25 17:18:57 -04:00
sshleifer 800dcdcce2 boom boom 2020-05-25 16:51:21 -04:00
sshleifer 88f9f917a6 Sortish Sampler 2020-05-25 16:44:48 -04:00
sshleifer a42f09eec6 warmup_steps=500 2020-05-25 16:24:49 -04:00
sshleifer c154fb3f8b style 2020-05-25 10:12:23 -04:00
sshleifer c7f6e62c3e boom boom 2020-05-25 09:26:59 -04:00
sshleifer af8b962edf boom boom 2020-05-25 09:26:11 -04:00
sshleifer 6f1757d2b8 boom boom 2020-05-25 09:23:48 -04:00
sshleifer 4d3617acc6 boom boom 2020-05-24 22:45:23 -04:00
sshleifer 2443a7eefc boom boom 2020-05-24 22:42:50 -04:00
sshleifer 6f03e39b30 boom boom 2020-05-24 22:39:17 -04:00
sshleifer 9f214624f4 boom boom 2020-05-24 22:37:28 -04:00
sshleifer e48a8299b9 boom boom 2020-05-24 22:36:48 -04:00
sshleifer 20abeec1f5 passing 2020-05-24 14:22:45 -04:00
sshleifer 584c329ffe encoder different 2020-05-24 14:20:29 -04:00
sshleifer d32887160a add RougeTracker 2020-05-22 13:02:59 -04:00
sshleifer ef45ee9a7e boom boom 2020-05-22 15:58:06 +00:00
sshleifer 8fbdd7d205 boom boom 2020-05-22 11:54:34 -04:00
sshleifer 1ca011ff75 test_mtl=140 default 2020-05-22 11:49:34 -04:00
sshleifer a78603cee7 boom boom 2020-05-21 23:14:56 -04:00
sshleifer d4830f9d0a rename 2020-05-21 23:13:57 -04:00
sshleifer c779c38297 boom boom 2020-05-21 17:04:19 -04:00
sshleifer 021d2b9f80 only train 2020-05-21 15:04:20 -04:00
sshleifer fa0eda43b6 boom boom 2020-05-21 14:19:55 -04:00
sshleifer c6f7e14d0b progbar 2020-05-21 14:07:16 -04:00
sshleifer db3d85ad2c boom boom 2020-05-21 13:56:35 -04:00
sshleifer 705ed2ffb4 Better desc 2020-05-21 13:43:47 -04:00
sshleifer 29499c9ca5 boom boom 2020-05-21 13:19:07 -04:00
sshleifer bcb9996ae5 overcautious freezing 2020-05-21 13:17:17 -04:00
sshleifer aa995f98a3 freeze after init 2020-05-21 13:02:00 -04:00
sshleifer 5d751729ec Freeze encoder fix 2020-05-21 12:55:59 -04:00
sshleifer f5606df47b boom boom 2020-05-21 12:26:07 -04:00
sshleifer ca82946559 boom boom 2020-05-21 12:24:06 -04:00
sshleifer 3461e24b8d boom boom 2020-05-21 12:21:42 -04:00
sshleifer 15a2d4ee79 boom boom 2020-05-21 11:45:07 -04:00
sshleifer 3f73cb76a7 boom boom 2020-05-21 11:41:13 -04:00
sshleifer 88b4b970d3 switch copy logic 2020-05-21 11:39:29 -04:00
sshleifer ac72c7aadc convenience method 2020-05-21 09:55:38 -04:00
sshleifer dd418e8d1d boom boom 2020-05-21 09:41:17 -04:00
sshleifer 07da0623ad avoid losing bart grouped batch sampler work 2020-05-21 08:49:24 -04:00
sshleifer 21a2cc0eb1 delete run_distiller -> main 2020-05-21 01:20:11 -04:00
sshleifer 2c2a10e5d7 orig tests pass 2020-05-21 01:08:43 -04:00
sshleifer 910da0fb96 boom boom 2020-05-21 01:00:57 -04:00
sshleifer 7a7dcd397f l2copy[1] = 0 2020-05-20 22:25:18 -04:00
sshleifer b39c39fe87 boom boom 2020-05-20 22:09:49 -04:00
sshleifer ce0de9073a boom boom 2020-05-20 22:07:35 -04:00
sshleifer cf47a6ff9d boom boom 2020-05-20 22:05:05 -04:00
sshleifer a6420767e9 boom boom 2020-05-20 22:04:48 -04:00
sshleifer fee71b4fef npars 2020-05-20 22:01:26 -04:00
sshleifer a2671a6eee boom boom 2020-05-20 21:30:49 -04:00
sshleifer 18524d8396 boom boom 2020-05-20 21:30:36 -04:00
sshleifer a250d1c6aa boom boom 2020-05-20 21:30:09 -04:00
sshleifer b62bfe666d spelling 2020-05-20 18:36:14 -04:00
sshleifer 1fc653f626 log_metrics 2020-05-20 18:18:11 -04:00
sshleifer f281724e92 wandb 2020-05-20 18:14:27 -04:00
sshleifer 67f9553d4b Run encoder once 2020-05-20 17:58:43 -04:00
sshleifer cd461e13ff fix 2020-05-20 16:40:27 -04:00
sshleifer 05bfb1eb1d Support fewer layers 2020-05-20 16:38:20 -04:00
sshleifer c79351a64a fix flatten 2020-05-20 15:59:22 -04:00
sshleifer c774160bd2 imports 2020-05-20 15:30:32 -04:00
sshleifer beda65dd17 bash 2020-05-20 19:29:21 +00:00
sshleifer 59c658de2f fixes 2020-05-20 15:28:05 -04:00
sshleifer 5947812374 add files 2020-05-20 15:08:41 -04:00
sshleifer 09cc1e64cf save preds 2020-05-20 15:04:38 -04:00
sshleifer 5fbf42258c summaries for file 2020-05-20 12:56:26 -04:00
sshleifer af315c157e boom boom 2020-05-20 09:55:53 -04:00
sshleifer 2c7294830e val check test passing with smaller batch size 2020-05-20 09:31:11 -04:00
sshleifer 2b7132c25e metrics saving, but no val_check_interval honored 2020-05-20 09:20:55 -04:00
sshleifer bef77fc211 boom boom 2020-05-19 15:59:22 -04:00
sshleifer abb81df15f boom boom 2020-05-19 15:58:04 -04:00
sshleifer 51221fbbf8 boom boom 2020-05-19 15:57:04 -04:00
sshleifer 2ee9388492 boom boom 2020-05-19 15:54:12 -04:00
sshleifer c4530f91ab batch 2020-05-19 15:22:13 -04:00
sshleifer bbc4e52d3b assert student small 2020-05-19 15:18:50 -04:00
sshleifer 74704cedab boom boom 2020-05-19 14:46:35 -04:00
sshleifer bf5782e080 boom boom 2020-05-19 14:38:41 -04:00
sshleifer 3cebd56848 boom boom 2020-05-19 14:35:05 -04:00
sshleifer 696c8c28f2 boom boom 2020-05-19 14:34:20 -04:00
sshleifer 04a8ace0df boom boom 2020-05-19 14:32:41 -04:00
sshleifer 0ed41562eb boom boom 2020-05-19 14:30:36 -04:00
sshleifer 7081d0f3fd add rouge 2020-05-19 14:29:14 -04:00
sshleifer ca9c685453 rouge 2020-05-19 14:05:46 -04:00
sshleifer 70cf536ae2 Merge branch 'distilbart' of github.com:sshleifer/transformers_fork into distilbart 2020-05-19 11:35:22 -04:00
sshleifer 5a3ed998f0 Cache tokenized 2020-05-19 11:35:17 -04:00
sshleifer 6302fb02ef bs=8 2020-05-19 14:44:21 +00:00
sshleifer 5a35811635 boom boom 2020-05-19 09:55:26 -04:00
sshleifer 174aaf374d Fast dev run 2020-05-19 09:54:43 -04:00
sshleifer 1edc50f670 bash 2020-05-19 13:50:01 +00:00
sshleifer d2cc12b182 boom boom 2020-05-19 09:16:02 -04:00
sshleifer 8afb88ce24 relatif importschlossen 2020-05-19 09:09:15 -04:00
sshleifer 4be1287671 real data test passes 2020-05-19 09:00:33 -04:00
sshleifer 4f5790f652 passing 2020-05-18 15:16:41 -04:00
sshleifer fcc49a0c5c can import 2020-05-18 14:31:36 -04:00
sshleifer 937d3d6ae8 Merge branch 'master' into distilbart 2020-05-18 14:30:34 -04:00
sshleifer b01f1d594c Failing test 2020-05-18 14:19:46 -04:00
sshleifer 71850ad26b copy decoder layers 2020-05-18 13:33:30 -04:00
sshleifer 1368969930 passing 2020-05-18 11:55:30 -04:00
sshleifer f2de528d04 passing 2020-05-18 11:49:12 -04:00
sshleifer eabb3c5f9a test passing, loss doesnt change 2020-05-18 11:13:54 -04:00
sshleifer 762dd832f1 boom boom 2020-05-17 15:43:35 -04:00
sshleifer 8280dad80e roberta extract test works 2020-05-16 11:09:26 -04:00
18 changed files with 1179 additions and 80 deletions
+13
View File
@@ -116,6 +116,19 @@ class BaseTransformer(pl.LightningModule):
self.opt = optimizer
return [optimizer]
def optimizer_step(self, epoch, batch_idx, optimizer, optimizer_idx, second_order_closure=None):
if self.trainer.use_tpu:
xm.optimizer_step(optimizer)
else:
optimizer.step()
optimizer.zero_grad()
self.lr_scheduler.step()
def get_tqdm_dict(self):
avg_loss = getattr(self.trainer, "avg_loss", 0.0)
tqdm_dict = {"loss": "{:.3f}".format(avg_loss), "lr": self.lr_scheduler.get_last_lr()[-1]}
return tqdm_dict
def test_step(self, batch, batch_nb):
return self.validation_step(batch, batch_nb)
+1 -1
View File
@@ -44,7 +44,7 @@ export me=`git config user.name`
```
Tips:
- 1 epoch at batch size 1 for bart-large takes 24 hours, requires 13GB GPU RAM with fp16 on an NVIDIA-V100.
- 1 epoch at batch size 1 for bart-large takes 24 hours, requires 13GB GPU RAM with fp16 on an NVIDIA-V100.
- try `bart-base`, `--freeze_encoder` or `--freeze_embeds` for faster training/larger batch size. (3hr/epoch with bs=8, see below)
- `fp16_opt_level=O1` (the default works best).
- If you are finetuning on your own dataset, start from `bart-large-cnn` if you want long summaries and `bart-large-xsum` if you want short summaries.
@@ -0,0 +1,103 @@
from pathlib import Path
import numpy as np
import pandas as pd
import torch
from durbango import lmap, tqdm_nice
from transformers import BartTokenizer
try:
from finetune import calculate_rouge
except ImportError:
from .finetune import calculate_rouge
LOGDIR = Path("examples/summarization/dbart/logs/").absolute()
DATA_DIR = Path("examples/summarization/dbart/cnn_dm").absolute()
def rouge_files(src_file: Path, tgt_file: Path):
src = lmap(str.strip, list(src_file.open().readlines()))
tgt = lmap(str.strip, list(tgt_file.open().readlines()))
return calculate_rouge(src, tgt)
def read_gens(exp_name, split="test", n=None):
if Path(exp_name).exists():
expdir = exp_name
else:
expdir = LOGDIR / exp_name
assert expdir.exists(), expdir
paths = list(expdir.glob(f"{split}_generations*.txt"))
assert paths
path = paths[0]
lns = lmap(str.strip, list(path.open().readlines()))
if n is not None:
return lns[:n]
return lns
def load_cpu(p):
return torch.load(p, map_location="cpu")
class RougeTracker:
def __init__(self, csv_path="rouge_test_df.csv", logdir=LOGDIR, data_dir=DATA_DIR):
try:
self.df = pd.read_csv(csv_path, index_col=0)
except FileNotFoundError:
self.df = pd.DataFrame()
self.logdir = logdir
test_gt = lmap(str.strip, Path(data_dir / "test.target").open().readlines())
test_gt = lmap(str.strip, test_gt)
self.gt = test_gt # {'test': test_gt}
self.tokenizer = BartTokenizer.from_pretrained("facebook/bart-large-cnn")
self.csv_path = csv_path
self.new_results = pd.DataFrame()
@property
def finished_experiments(self):
return [p.parent.name for p in list(self.logdir.glob("*/test_generations*.txt"))]
def tok_len(self, strang):
return len(self.tokenizer.encode(strang))
def read_all_gens(self,):
GENS = {}
for f in self.finished_experiments:
GENS[f] = read_gens(f)
return GENS
def score(self, gens, k, all_stats=False):
rouge_raw = calculate_rouge(gens, self.gt, all_stats=all_stats)
lens = np.mean(lmap(self.tok_len, gens))
return dict(avg_len=lens, exp_name=k, **rouge_raw)
@property
def new_experiments(self):
possible = set(self.finished_experiments).difference(self.df.index)
to_score = set()
for p in possible:
gens = read_gens(p)
if len(gens) == len(self.gt):
to_score.add(p)
return to_score
def update(self):
records = []
to_score = self.new_experiments
for exp_name in tqdm_nice(to_score, desc="Rouge Update"):
gens = read_gens(exp_name)
if len(gens) != len(self.gt):
continue
records.append(self.score(gens, exp_name))
if not records:
return self.df
new_df = (
pd.DataFrame(records).rename(columns=lambda x: x.replace("rouge", "R")).set_index("exp_name").astype(float)
)
self.new_results = new_df
self.df = pd.concat([self.df, new_df]).dsort("R2")
return self.df
+102 -32
View File
@@ -25,6 +25,8 @@ try:
any_requires_grad,
)
from .finetune import main as ft_main
from .replacement_scheduler import LinearReplacementScheduler
except ImportError:
from finetune import SummarizationModule
from finetune import main as ft_main
@@ -37,6 +39,62 @@ except ImportError:
assert_all_frozen,
any_requires_grad,
)
from replacement_scheduler import LinearReplacementScheduler
class TheseusDistiller(SummarizationModule):
def __init__(self, hparams):
assert Path(hparams.data_dir).exists()
# config = BartConfig.from_pretrained(hparams.model_name_or_path, student_encoder_layers=hparams.student_encoder_layers, student_decoder_layers=hparams.student_decoder_layers, replacing_rate=hparams.theseus_replace_rate)
model = BartForConditionalGeneration.from_pretrained(
hparams.model_name_or_path,
student_encoder_layers=hparams.student_encoder_layers,
student_decoder_layers=hparams.student_decoder_layers,
replacing_rate=hparams.theseus_replace_rate,
)
super().__init__(hparams, model=model)
self.different_encoder: bool = hparams.student_encoder_layers != self.model.config.encoder_layers
self.different_decoder: bool = hparams.student_decoder_layers != self.model.config.decoder_layers
if hparams.theseus_init_copy:
hparams.d_layers_to_copy = get_layers_to_copy(
hparams.student_decoder_layers, self.model.config.decoder_layers, strategy=hparams.init_strategy,
)
hparams.e_layers_to_copy: List = get_layers_to_copy(
hparams.student_encoder_layers, self.model.config.encoder_layers, strategy=hparams.init_strategy
)
if self.different_decoder:
copy_layers(
self.model.model.decoder.layers, self.model.model.decoder.scc_layers, hparams.d_layers_to_copy
)
if self.different_encoder:
copy_layers(
self.model.model.encoder.layers, self.model.model.encoder.scc_layers, hparams.e_layers_to_copy
)
else:
hparams.e_layers_to_copy, hparams.d_layers_to_copy = None, None
self.replace_scheduler_encoder = LinearReplacementScheduler(self.model.model.encoder, 0.6)
self.replace_scheduler_decoder = LinearReplacementScheduler(self.model.model.decoder, 0.6)
freeze_params(self.model.model.encoder.layers) # Test
freeze_params(self.model.model.decoder.layers)
def optimizer_step(self, *args, **kwargs) -> None:
self.replace_scheduler_encoder.step()
replace_rate = self.replace_scheduler_decoder.step()
self.logger.log_metrics({"replace_rate": replace_rate})
super().optimizer_step(*args, **kwargs)
def copy_to_student(self, d_layers_to_copy, e_layers_to_copy, hparams, student, teacher):
if teacher.config.model_type == "t5":
return self.copy_t5_to_student(d_layers_to_copy, e_layers_to_copy, hparams, student, teacher)
self.different_encoder: bool = hparams.student_encoder_layers != teacher.config.encoder_layers
self.different_decoder = hparams.student_decoder_layers != teacher.config.decoder_layers
if self.different_decoder:
copy_layers(teacher.model.decoder.layers, student.model.decoder.layers, d_layers_to_copy)
if self.different_encoder:
copy_layers(teacher.model.encoder.layers, student.model.encoder.layers, e_layers_to_copy)
class SummarizationDistiller(SummarizationModule):
@@ -49,6 +107,9 @@ class SummarizationDistiller(SummarizationModule):
super().__init__(hparams, model=student, config=student_cfg)
self.teacher = teacher
if isinstance(self.teacher, BartForConditionalGeneration):
assert teacher.model.encoder.scc_layers is None
assert self.model.model.encoder.scc_layers is None
use_task_specific_params(self.teacher, "summarization")
freeze_params(self.teacher)
self.sanity_check_gradients()
@@ -79,8 +140,14 @@ class SummarizationDistiller(SummarizationModule):
"decoder_layers": hparams.student_decoder_layers,
"encoder_layers": hparams.student_encoder_layers,
}
d_layers_to_copy = get_layers_to_copy(student_updates["decoder_layers"], teacher.config.decoder_layers)
e_layers_to_copy: List = get_layers_to_copy(student_updates["encoder_layers"], teacher.config.encoder_layers)
d_layers_to_copy = get_layers_to_copy(
student_updates["decoder_layers"], teacher.config.decoder_layers, strategy=hparams.init_strategy
)
e_layers_to_copy: List = get_layers_to_copy(
student_updates["encoder_layers"], teacher.config.encoder_layers, strategy=hparams.init_strategy
)
hparams.d_layer_to_copy = d_layers_to_copy
hparams.e_layer_to_copy = e_layers_to_copy
kw = teacher.config.to_diff_dict()
@@ -180,18 +247,13 @@ class SummarizationDistiller(SummarizationModule):
# parser.add_argument("--alpha_cos", default=0.0, type=float)
parser.add_argument("--alpha_encoder_loss", default=0.0, type=float)
parser.add_argument("--alpha_hid", default=0.0, type=float, required=False)
parser.add_argument(
"--student_decoder_layers", default=12, type=int, required=False,
)
parser.add_argument(
"--student_encoder_layers", default=12, type=int, required=False,
)
parser.add_argument(
"--no_teacher", action="store_true", default=False,
)
parser.add_argument( # TODO: remove
"--enc_only", action="store_true", default=False,
)
parser.add_argument("--student_decoder_layers", default=12, type=int, required=False)
parser.add_argument("--student_encoder_layers", default=12, type=int, required=False)
parser.add_argument("--no_teacher", action="store_true", default=False)
parser.add_argument("--theseus_replace_rate", type=float, default=0.0)
parser.add_argument("--theseus_init_copy", action="store_true")
parser.add_argument("--init_strategy", type=str, default="alternate", choices=["alternate", "top", "bottom"])
return parser
def _step(self, batch):
@@ -378,13 +440,12 @@ class T5SummarizationDistiller(SummarizationDistiller):
def create_module(args):
t5 = "t5" in args.model_name_or_path
if args.no_teacher:
assert not args.enc_only
if args.no_teacher and args.theseus_replace_rate == 0:
module_cls = SummarizationModule
elif args.no_teacher and args.theseus_replace_rate > 0:
module_cls = TheseusDistiller
elif t5:
module_cls = T5SummarizationDistiller
elif args.enc_only:
raise ValueError("Deleted that")
else:
module_cls = SummarizationDistiller
args.setup_cls: str = module_cls.__name__
@@ -415,21 +476,30 @@ def evaluate_checkpoint(ckpt_path: Path, dest_dir=None):
trainer.test(model)
def get_layers_to_copy(n_to_get, tot):
all_layers = list(range(tot))
if tot == 12: # Alternating for special cases
layers_to_copy = { # maps # layers in student -> which teacher layers to copy
6: [0, 2, 4, 7, 9, 11],
1: [11],
3: [0, 6, 11],
2: [0, 11],
4: [0, 4, 8, 11],
9: [0, 1, 2, 4, 5, 7, 9, 10, 11],
12: all_layers,
}
return layers_to_copy[n_to_get]
DISTILBERT_ALTERNATE_PATTERN = { # maps # layers in student -> which teacher layers to copy
6: [0, 2, 4, 7, 9, 11],
1: [11],
3: [0, 6, 11],
2: [0, 11],
4: [0, 4, 8, 11],
9: [0, 1, 2, 4, 5, 7, 9, 10, 11],
12: list(range(12)),
}
def get_layers_to_copy(n_student_layers: int, n_teacher_layers: int, strategy="alternate") -> List:
all_layers = list(range(n_teacher_layers))
if strategy == "alternate":
if n_teacher_layers == 12:
return DISTILBERT_ALTERNATE_PATTERN[n_student_layers]
else:
return all_layers[::2][:n_student_layers]
elif strategy == "bottom":
return all_layers[:n_student_layers]
elif strategy == "top":
return all_layers[-n_student_layers:]
else:
return all_layers[:n_to_get]
raise ValueError(f"layer copy strategy {strategy} not supported")
def distill_main(args):
+10 -4
View File
@@ -9,12 +9,14 @@ from typing import Dict, List, Tuple
import numpy as np
import pytorch_lightning as pl
import torch
from pytorch_lightning.loggers import WandbLogger
from torch.utils.data import DataLoader
from lightning_base import BaseTransformer, add_generic_args, generic_train
from transformers import get_linear_schedule_with_warmup
WANDB_PROJ_NAME = "transformers_fork-examples_summarization_bart"
try:
from .utils import (
use_task_specific_params,
@@ -52,6 +54,7 @@ class SummarizationModule(BaseTransformer):
loss_names = ["loss"]
def __init__(self, hparams, **kwargs):
assert Path(hparams.data_dir).exists()
super().__init__(hparams, num_labels=None, mode=self.mode, **kwargs)
use_task_specific_params(self.model, "summarization")
save_git_info(self.hparams.output_dir)
@@ -79,8 +82,7 @@ class SummarizationModule(BaseTransformer):
}
assert self.target_lens["train"] <= self.target_lens["val"], f"target_lens: {self.target_lens}"
assert self.target_lens["train"] <= self.target_lens["test"], f"target_lens: {self.target_lens}"
if self.hparams.freeze_embeds:
if not self.hparams.unfreeze_embeds:
self.freeze_embeds()
if self.hparams.freeze_encoder:
freeze_params(self.model.model.encoder) # TODO: this will break for t5
@@ -253,9 +255,9 @@ class SummarizationModule(BaseTransformer):
help="The input data dir. Should contain train.source, train.target, val.source, val.target, test.source, test.target",
)
parser.add_argument("--freeze_encoder", action="store_true")
parser.add_argument("--freeze_embeds", action="store_true")
parser.add_argument("--unfreeze_embeds", action="store_true")
parser.add_argument("--sortish_sampler", action="store_true", default=False)
parser.add_argument("--logger", type=str, choices=["default", "wandb", "wandb_shared"], default="default")
parser.add_argument("--logger", type=str, choices=["default", "wandb", "wandb_shared"], default="wandb")
parser.add_argument("--n_train", type=int, default=-1, required=False, help="# examples. -1 means use all.")
parser.add_argument("--n_val", type=int, default=500, required=False, help="# examples. -1 means use all.")
parser.add_argument("--n_test", type=int, default=-1, required=False, help="# examples. -1 means use all.")
@@ -278,6 +280,10 @@ def main(args, model=None) -> SummarizationModule:
elif args.logger == "wandb":
from pytorch_lightning.loggers import WandbLogger
logger = WandbLogger(name=model.output_dir.name, project=WANDB_PROJ_NAME)
elif args.logger == "wandb_shared":
from pytorch_lightning.loggers import WandbLogger
logger = WandbLogger(name=model.output_dir.name)
elif args.logger == "wandb_shared":
from pytorch_lightning.loggers import WandbLogger
+1 -1
View File
@@ -1,4 +1,3 @@
# Add parent directory to python path to access lightning_base.py
export PYTHONPATH="../":"${PYTHONPATH}"
@@ -6,6 +5,7 @@ export PYTHONPATH="../":"${PYTHONPATH}"
# --model_name_or_path=t5-base for t5
# the proper usage is documented in the README
python finetune.py \
--model_name_or_path=facebook/bart-large \
--learning_rate=3e-5 \
@@ -0,0 +1,639 @@
"""PyTorch BERT-of-Theseus model. """
from __future__ import absolute_import, division, print_function, unicode_literals
import logging
import torch
from torch import nn
from torch.distributions.bernoulli import Bernoulli
from torch.nn import CrossEntropyLoss, MSELoss
from transformers.configuration_bert import BertConfig
from transformers.modeling_bert import (
ACT2FN,
BERT_PRETRAINED_MODEL_ARCHIVE_MAP,
BertAttention,
BertEmbeddings,
BertIntermediate,
BertLayer,
BertLayerNorm,
BertLMPredictionHead,
BertOnlyMLMHead,
BertOnlyNSPHead,
BertOutput,
BertPooler,
BertPredictionHeadTransform,
BertPreTrainingHeads,
BertSelfAttention,
BertSelfOutput,
gelu,
gelu_new,
load_tf_weights_in_bert,
mish,
swish,
)
from transformers.modeling_utils import PreTrainedModel, prune_linear_layer
logger = logging.getLogger(__name__)
class BertEncoder(nn.Module):
def __init__(self, config, scc_n_layer=6):
super(BertEncoder, self).__init__()
self.prd_n_layer = config.num_hidden_layers
self.scc_n_layer = scc_n_layer
assert self.prd_n_layer % self.scc_n_layer == 0
self.compress_ratio = self.prd_n_layer // self.scc_n_layer
self.bernoulli = None
self.output_attentions = config.output_attentions
self.output_hidden_states = config.output_hidden_states
self.layer = nn.ModuleList([BertLayer(config) for _ in range(self.prd_n_layer)])
self.scc_layer = nn.ModuleList([BertLayer(config) for _ in range(self.scc_n_layer)])
def set_replacing_rate(self, replacing_rate):
if not 0 < replacing_rate <= 1:
raise Exception("Replace rate must be in the range (0, 1]!")
self.bernoulli = Bernoulli(torch.tensor([replacing_rate]))
def forward(
self,
hidden_states,
attention_mask=None,
head_mask=None,
encoder_hidden_states=None,
encoder_attention_mask=None,
):
all_hidden_states = ()
all_attentions = ()
if self.training:
inference_layers = []
for i in range(self.scc_n_layer):
if self.bernoulli.sample() == 1: # REPLACE
inference_layers.append(self.scc_layer[i])
else: # KEEP the original
for offset in range(self.compress_ratio):
inference_layers.append(self.layer[i * self.compress_ratio + offset])
else: # inference with compressed model
inference_layers = self.scc_layer
for i, layer_module in enumerate(inference_layers):
if self.output_hidden_states:
all_hidden_states = all_hidden_states + (hidden_states,)
layer_outputs = layer_module(
hidden_states, attention_mask, head_mask[i], encoder_hidden_states, encoder_attention_mask
)
hidden_states = layer_outputs[0]
if self.output_attentions:
all_attentions = all_attentions + (layer_outputs[1],)
# Add last layer
if self.output_hidden_states:
all_hidden_states = all_hidden_states + (hidden_states,)
outputs = (hidden_states,)
if self.output_hidden_states:
outputs = outputs + (all_hidden_states,)
if self.output_attentions:
outputs = outputs + (all_attentions,)
return outputs # last-layer hidden state, (all hidden states), (all attentions)
class BertPreTrainedModel(PreTrainedModel):
config_class = BertConfig
pretrained_model_archive_map = BERT_PRETRAINED_MODEL_ARCHIVE_MAP
load_tf_weights = load_tf_weights_in_bert
base_model_prefix = "bert"
def _init_weights(self, module):
""" Initialize the weights """
if isinstance(module, (nn.Linear, nn.Embedding)):
# Slightly different from the TF version which uses truncated_normal for initialization
# cf https://github.com/pytorch/pytorch/pull/5617
module.weight.data.normal_(mean=0.0, std=self.config.initializer_range)
elif isinstance(module, BertLayerNorm):
module.bias.data.zero_()
module.weight.data.fill_(1.0)
if isinstance(module, nn.Linear) and module.bias is not None:
module.bias.data.zero_()
class BertModel(BertPreTrainedModel):
def __init__(self, config):
super(BertModel, self).__init__(config)
self.config = config
self.embeddings = BertEmbeddings(config)
self.encoder = BertEncoder(config)
self.pooler = BertPooler(config)
self.init_weights()
def get_input_embeddings(self):
return self.embeddings.word_embeddings
def set_input_embeddings(self, value):
self.embeddings.word_embeddings = value
def _prune_heads(self, heads_to_prune):
""" Prunes heads of the model.
heads_to_prune: dict of {layer_num: list of heads to prune in this layer}
See base class PreTrainedModel
"""
for layer, heads in heads_to_prune.items():
self.encoder.layer[layer].attention.prune_heads(heads)
def forward(
self,
input_ids=None,
attention_mask=None,
token_type_ids=None,
position_ids=None,
head_mask=None,
inputs_embeds=None,
encoder_hidden_states=None,
encoder_attention_mask=None,
):
if input_ids is not None and inputs_embeds is not None:
raise ValueError("You cannot specify both input_ids and inputs_embeds at the same time")
elif input_ids is not None:
input_shape = input_ids.size()
elif inputs_embeds is not None:
input_shape = inputs_embeds.size()[:-1]
else:
raise ValueError("You have to specify either input_ids or inputs_embeds")
device = input_ids.device if input_ids is not None else inputs_embeds.device
if attention_mask is None:
attention_mask = torch.ones(input_shape, device=device)
if token_type_ids is None:
token_type_ids = torch.zeros(input_shape, dtype=torch.long, device=device)
# We can provide a self-attention mask of dimensions [batch_size, from_seq_length, to_seq_length]
# ourselves in which case we just need to make it broadcastable to all heads.
if attention_mask.dim() == 3:
extended_attention_mask = attention_mask[:, None, :, :]
elif attention_mask.dim() == 2:
# Provided a padding mask of dimensions [batch_size, seq_length]
# - if the model is a decoder, apply a causal mask in addition to the padding mask
# - if the model is an encoder, make the mask broadcastable to [batch_size, num_heads, seq_length, seq_length]
if self.config.is_decoder:
batch_size, seq_length = input_shape
seq_ids = torch.arange(seq_length, device=device)
causal_mask = seq_ids[None, None, :].repeat(batch_size, seq_length, 1) <= seq_ids[None, :, None]
causal_mask = causal_mask.to(
torch.long
) # not converting to long will cause errors with pytorch version < 1.3
extended_attention_mask = causal_mask[:, None, :, :] * attention_mask[:, None, None, :]
else:
extended_attention_mask = attention_mask[:, None, None, :]
else:
raise ValueError(
"Wrong shape for input_ids (shape {}) or attention_mask (shape {})".format(
input_shape, attention_mask.shape
)
)
# Since attention_mask is 1.0 for positions we want to attend and 0.0 for
# masked positions, this operation will create a tensor which is 0.0 for
# positions we want to attend and -10000.0 for masked positions.
# Since we are adding it to the raw scores before the softmax, this is
# effectively the same as removing these entirely.
extended_attention_mask = extended_attention_mask.to(dtype=next(self.parameters()).dtype) # fp16 compatibility
extended_attention_mask = (1.0 - extended_attention_mask) * -10000.0
# If a 2D ou 3D attention mask is provided for the cross-attention
# we need to make broadcastabe to [batch_size, num_heads, seq_length, seq_length]
if self.config.is_decoder and encoder_hidden_states is not None:
encoder_batch_size, encoder_sequence_length, _ = encoder_hidden_states.size()
encoder_hidden_shape = (encoder_batch_size, encoder_sequence_length)
if encoder_attention_mask is None:
encoder_attention_mask = torch.ones(encoder_hidden_shape, device=device)
if encoder_attention_mask.dim() == 3:
encoder_extended_attention_mask = encoder_attention_mask[:, None, :, :]
elif encoder_attention_mask.dim() == 2:
encoder_extended_attention_mask = encoder_attention_mask[:, None, None, :]
else:
raise ValueError(
"Wrong shape for encoder_hidden_shape (shape {}) or encoder_attention_mask (shape {})".format(
encoder_hidden_shape, encoder_attention_mask.shape
)
)
encoder_extended_attention_mask = encoder_extended_attention_mask.to(
dtype=next(self.parameters()).dtype
) # fp16 compatibility
encoder_extended_attention_mask = (1.0 - encoder_extended_attention_mask) * -10000.0
else:
encoder_extended_attention_mask = None
# Prepare head mask if needed
# 1.0 in head_mask indicate we keep the head
# attention_probs has shape bsz x n_heads x N x N
# input head_mask has shape [num_heads] or [num_hidden_layers x num_heads]
# and head_mask is converted to shape [num_hidden_layers x batch x num_heads x seq_length x seq_length]
if head_mask is not None:
if head_mask.dim() == 1:
head_mask = head_mask.unsqueeze(0).unsqueeze(0).unsqueeze(-1).unsqueeze(-1)
head_mask = head_mask.expand(self.config.num_hidden_layers, -1, -1, -1, -1)
elif head_mask.dim() == 2:
head_mask = (
head_mask.unsqueeze(1).unsqueeze(-1).unsqueeze(-1)
) # We can specify head_mask for each layer
head_mask = head_mask.to(
dtype=next(self.parameters()).dtype
) # switch to fload if need + fp16 compatibility
else:
head_mask = [None] * self.config.num_hidden_layers
embedding_output = self.embeddings(
input_ids=input_ids, position_ids=position_ids, token_type_ids=token_type_ids, inputs_embeds=inputs_embeds
)
encoder_outputs = self.encoder(
embedding_output,
attention_mask=extended_attention_mask,
head_mask=head_mask,
encoder_hidden_states=encoder_hidden_states,
encoder_attention_mask=encoder_extended_attention_mask,
)
sequence_output = encoder_outputs[0]
pooled_output = self.pooler(sequence_output)
outputs = (sequence_output, pooled_output,) + encoder_outputs[
1:
] # add hidden_states and attentions if they are here
return outputs # sequence_output, pooled_output, (hidden_states), (attentions)
class BertForPreTraining(BertPreTrainedModel):
def __init__(self, config):
super(BertForPreTraining, self).__init__(config)
self.bert = BertModel(config)
self.cls = BertPreTrainingHeads(config)
self.init_weights()
def get_output_embeddings(self):
return self.cls.predictions.decoder
def forward(
self,
input_ids=None,
attention_mask=None,
token_type_ids=None,
position_ids=None,
head_mask=None,
inputs_embeds=None,
masked_lm_labels=None,
next_sentence_label=None,
):
outputs = self.bert(
input_ids,
attention_mask=attention_mask,
token_type_ids=token_type_ids,
position_ids=position_ids,
head_mask=head_mask,
inputs_embeds=inputs_embeds,
)
sequence_output, pooled_output = outputs[:2]
prediction_scores, seq_relationship_score = self.cls(sequence_output, pooled_output)
outputs = (prediction_scores, seq_relationship_score,) + outputs[
2:
] # add hidden states and attention if they are here
if masked_lm_labels is not None and next_sentence_label is not None:
loss_fct = CrossEntropyLoss(ignore_index=-1)
masked_lm_loss = loss_fct(prediction_scores.view(-1, self.config.vocab_size), masked_lm_labels.view(-1))
next_sentence_loss = loss_fct(seq_relationship_score.view(-1, 2), next_sentence_label.view(-1))
total_loss = masked_lm_loss + next_sentence_loss
outputs = (total_loss,) + outputs
return outputs # (loss), prediction_scores, seq_relationship_score, (hidden_states), (attentions)
class BertForMaskedLM(BertPreTrainedModel):
def __init__(self, config):
super(BertForMaskedLM, self).__init__(config)
self.bert = BertModel(config)
self.cls = BertOnlyMLMHead(config)
self.init_weights()
def get_output_embeddings(self):
return self.cls.predictions.decoder
def forward(
self,
input_ids=None,
attention_mask=None,
token_type_ids=None,
position_ids=None,
head_mask=None,
inputs_embeds=None,
masked_lm_labels=None,
encoder_hidden_states=None,
encoder_attention_mask=None,
lm_labels=None,
):
outputs = self.bert(
input_ids,
attention_mask=attention_mask,
token_type_ids=token_type_ids,
position_ids=position_ids,
head_mask=head_mask,
inputs_embeds=inputs_embeds,
encoder_hidden_states=encoder_hidden_states,
encoder_attention_mask=encoder_attention_mask,
)
sequence_output = outputs[0]
prediction_scores = self.cls(sequence_output)
outputs = (prediction_scores,) + outputs[2:] # Add hidden states and attention if they are here
# Although this may seem awkward, BertForMaskedLM supports two scenarios:
# 1. If a tensor that contains the indices of masked labels is provided,
# the cross-entropy is the MLM cross-entropy that measures the likelihood
# of predictions for masked words.
# 2. If `lm_labels` is provided we are in a causal scenario where we
# try to predict the next token for each input in the decoder.
if masked_lm_labels is not None:
loss_fct = CrossEntropyLoss(ignore_index=-1) # -1 index = padding token
masked_lm_loss = loss_fct(prediction_scores.view(-1, self.config.vocab_size), masked_lm_labels.view(-1))
outputs = (masked_lm_loss,) + outputs
if lm_labels is not None:
# we are doing next-token prediction; shift prediction scores and input ids by one
prediction_scores = prediction_scores[:, :-1, :].contiguous()
lm_labels = lm_labels[:, 1:].contiguous()
loss_fct = CrossEntropyLoss(ignore_index=-1)
ltr_lm_loss = loss_fct(prediction_scores.view(-1, self.config.vocab_size), lm_labels.view(-1))
outputs = (ltr_lm_loss,) + outputs
return outputs # (masked_lm_loss), (ltr_lm_loss), prediction_scores, (hidden_states), (attentions)
class BertForNextSentencePrediction(BertPreTrainedModel):
def __init__(self, config):
super(BertForNextSentencePrediction, self).__init__(config)
self.bert = BertModel(config)
self.cls = BertOnlyNSPHead(config)
self.init_weights()
def forward(
self,
input_ids=None,
attention_mask=None,
token_type_ids=None,
position_ids=None,
head_mask=None,
inputs_embeds=None,
next_sentence_label=None,
):
outputs = self.bert(
input_ids,
attention_mask=attention_mask,
token_type_ids=token_type_ids,
position_ids=position_ids,
head_mask=head_mask,
inputs_embeds=inputs_embeds,
)
pooled_output = outputs[1]
seq_relationship_score = self.cls(pooled_output)
outputs = (seq_relationship_score,) + outputs[2:] # add hidden states and attention if they are here
if next_sentence_label is not None:
loss_fct = CrossEntropyLoss(ignore_index=-1)
next_sentence_loss = loss_fct(seq_relationship_score.view(-1, 2), next_sentence_label.view(-1))
outputs = (next_sentence_loss,) + outputs
return outputs # (next_sentence_loss), seq_relationship_score, (hidden_states), (attentions)
class BertForSequenceClassification(BertPreTrainedModel):
def __init__(self, config):
super(BertForSequenceClassification, self).__init__(config)
self.num_labels = config.num_labels
self.bert = BertModel(config)
self.dropout = nn.Dropout(config.hidden_dropout_prob)
self.classifier = nn.Linear(config.hidden_size, self.config.num_labels)
self.init_weights()
def forward(
self,
input_ids=None,
attention_mask=None,
token_type_ids=None,
position_ids=None,
head_mask=None,
inputs_embeds=None,
labels=None,
):
outputs = self.bert(
input_ids,
attention_mask=attention_mask,
token_type_ids=token_type_ids,
position_ids=position_ids,
head_mask=head_mask,
inputs_embeds=inputs_embeds,
)
pooled_output = outputs[1]
pooled_output = self.dropout(pooled_output)
logits = self.classifier(pooled_output)
outputs = (logits,) + outputs[2:] # add hidden states and attention if they are here
if labels is not None:
if self.num_labels == 1:
# We are doing regression
loss_fct = MSELoss()
loss = loss_fct(logits.view(-1), labels.view(-1))
else:
loss_fct = CrossEntropyLoss()
loss = loss_fct(logits.view(-1, self.num_labels), labels.view(-1))
outputs = (loss,) + outputs
return outputs # (loss), logits, (hidden_states), (attentions)
class BertForMultipleChoice(BertPreTrainedModel):
def __init__(self, config):
super(BertForMultipleChoice, self).__init__(config)
self.bert = BertModel(config)
self.dropout = nn.Dropout(config.hidden_dropout_prob)
self.classifier = nn.Linear(config.hidden_size, 1)
self.init_weights()
def forward(
self,
input_ids=None,
attention_mask=None,
token_type_ids=None,
position_ids=None,
head_mask=None,
inputs_embeds=None,
labels=None,
):
num_choices = input_ids.shape[1]
input_ids = input_ids.view(-1, input_ids.size(-1))
attention_mask = attention_mask.view(-1, attention_mask.size(-1)) if attention_mask is not None else None
token_type_ids = token_type_ids.view(-1, token_type_ids.size(-1)) if token_type_ids is not None else None
position_ids = position_ids.view(-1, position_ids.size(-1)) if position_ids is not None else None
outputs = self.bert(
input_ids,
attention_mask=attention_mask,
token_type_ids=token_type_ids,
position_ids=position_ids,
head_mask=head_mask,
inputs_embeds=inputs_embeds,
)
pooled_output = outputs[1]
pooled_output = self.dropout(pooled_output)
logits = self.classifier(pooled_output)
reshaped_logits = logits.view(-1, num_choices)
outputs = (reshaped_logits,) + outputs[2:] # add hidden states and attention if they are here
if labels is not None:
loss_fct = CrossEntropyLoss()
loss = loss_fct(reshaped_logits, labels)
outputs = (loss,) + outputs
return outputs # (loss), reshaped_logits, (hidden_states), (attentions)
class BertForTokenClassification(BertPreTrainedModel):
def __init__(self, config):
super(BertForTokenClassification, self).__init__(config)
self.num_labels = config.num_labels
self.bert = BertModel(config)
self.dropout = nn.Dropout(config.hidden_dropout_prob)
self.classifier = nn.Linear(config.hidden_size, config.num_labels)
self.init_weights()
def forward(
self,
input_ids=None,
attention_mask=None,
token_type_ids=None,
position_ids=None,
head_mask=None,
inputs_embeds=None,
labels=None,
):
outputs = self.bert(
input_ids,
attention_mask=attention_mask,
token_type_ids=token_type_ids,
position_ids=position_ids,
head_mask=head_mask,
inputs_embeds=inputs_embeds,
)
sequence_output = outputs[0]
sequence_output = self.dropout(sequence_output)
logits = self.classifier(sequence_output)
outputs = (logits,) + outputs[2:] # add hidden states and attention if they are here
if labels is not None:
loss_fct = CrossEntropyLoss()
# Only keep active parts of the loss
if attention_mask is not None:
active_loss = attention_mask.view(-1) == 1
active_logits = logits.view(-1, self.num_labels)[active_loss]
active_labels = labels.view(-1)[active_loss]
loss = loss_fct(active_logits, active_labels)
else:
loss = loss_fct(logits.view(-1, self.num_labels), labels.view(-1))
outputs = (loss,) + outputs
return outputs # (loss), scores, (hidden_states), (attentions)
class BertForQuestionAnswering(BertPreTrainedModel):
def __init__(self, config):
super(BertForQuestionAnswering, self).__init__(config)
self.num_labels = config.num_labels
self.bert = BertModel(config)
self.qa_outputs = nn.Linear(config.hidden_size, config.num_labels)
self.init_weights()
def forward(
self,
input_ids=None,
attention_mask=None,
token_type_ids=None,
position_ids=None,
head_mask=None,
inputs_embeds=None,
start_positions=None,
end_positions=None,
):
outputs = self.bert(
input_ids,
attention_mask=attention_mask,
token_type_ids=token_type_ids,
position_ids=position_ids,
head_mask=head_mask,
inputs_embeds=inputs_embeds,
)
sequence_output = outputs[0]
logits = self.qa_outputs(sequence_output)
start_logits, end_logits = logits.split(1, dim=-1)
start_logits = start_logits.squeeze(-1)
end_logits = end_logits.squeeze(-1)
outputs = (start_logits, end_logits,) + outputs[2:]
if start_positions is not None and end_positions is not None:
# If we are on multi-GPU, split add a dimension
if len(start_positions.size()) > 1:
start_positions = start_positions.squeeze(-1)
if len(end_positions.size()) > 1:
end_positions = end_positions.squeeze(-1)
# sometimes the start/end positions are outside our model inputs, we ignore these terms
ignored_index = start_logits.size(1)
start_positions.clamp_(0, ignored_index)
end_positions.clamp_(0, ignored_index)
loss_fct = CrossEntropyLoss(ignore_index=ignored_index)
start_loss = loss_fct(start_logits, start_positions)
end_loss = loss_fct(end_logits, end_positions)
total_loss = (start_loss + end_loss) / 2
outputs = (total_loss,) + outputs
return outputs # (loss), start_logits, end_logits, (hidden_states), (attentions)
@@ -0,0 +1,32 @@
class ConstantReplacementScheduler:
def __init__(self, module, replacing_rate, replacing_steps=None):
self.module = module
self.replacing_rate = replacing_rate
self.replacing_steps = replacing_steps
self.step_counter = 0
self.module.set_replacing_rate(replacing_rate)
def step(self):
self.step_counter += 1
if self.replacing_steps is None or self.replacing_rate == 1.0:
return self.replacing_rate
else:
if self.step_counter >= self.replacing_steps:
self.module.set_replacing_rate(1.0)
self.replacing_rate = 1.0
return self.replacing_rate
class LinearReplacementScheduler:
def __init__(self, module, base_replacing_rate, k=1e-4):
self.module = module
self.base_replacing_rate = base_replacing_rate
self.step_counter = 0
self.k = k
self.module.set_replacing_rate(base_replacing_rate)
def step(self):
self.step_counter += 1
current_replacing_rate = min(self.k * self.step_counter + self.base_replacing_rate, 1.0)
self.module.set_replacing_rate(current_replacing_rate)
return current_replacing_rate
+1
View File
@@ -7,5 +7,6 @@ python distillation.py \
--learning_rate=3e-4 \
--do_train \
--do_predict \
--fp16 \
--val_check_interval 0.1 \
$@
+26 -17
View File
@@ -5,7 +5,7 @@ from pathlib import Path
import torch
from tqdm import tqdm
from transformers import AutoModelForSeq2SeqLM, AutoTokenizer
from transformers import AutoModelWithLMHead, AutoTokenizer
try:
@@ -23,13 +23,10 @@ def chunks(lst, n):
def generate_summaries(
examples: list, out_file: str, model_name: str, batch_size: int = 8, device: str = DEFAULT_DEVICE, fp16=False,
) -> None:
examples: list, out_file: str, model_name: str, batch_size: int = 8, device: str = DEFAULT_DEVICE
):
fout = Path(out_file).open("w", encoding="utf-8")
model_name = str(model_name)
model = AutoModelForSeq2SeqLM.from_pretrained(model_name).to(device)
if fp16:
model = model.half()
model = AutoModelWithLMHead.from_pretrained(model_name).to(device)
tokenizer = AutoTokenizer.from_pretrained(model_name)
@@ -52,20 +49,32 @@ def generate_summaries(
def run_generate():
parser = argparse.ArgumentParser()
parser.add_argument("input_path", type=str, help="like cnn_dm/test.source")
parser.add_argument("output_path", type=str, help="where to save summaries")
parser.add_argument("model_name", type=str, help="like facebook/bart-large-cnn,t5-base, etc.")
parser.add_argument(
"input_path", type=str, help="like cnn_dm/test.source",
)
parser.add_argument(
"output_path", type=str, help="where to save summaries",
)
parser.add_argument(
"model_name",
type=str,
default="facebook/bart-large-cnn",
help="like facebook/bart-large-cnn,'t5-small', 't5-base', 't5-large', 't5-3b', 't5-11b",
)
parser.add_argument("--reference_path", type=str, required=False, help="like cnn_dm/test_reference_summaries.txt")
parser.add_argument("--score_path", type=str, required=False, help="where to save the rouge score in json format")
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")
parser.add_argument("--fp16", action="store_true")
parser.add_argument(
"--score_path", type=str, required=False, help="where to save the rouge score in json format",
)
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",
)
args = parser.parse_args()
examples = [" " + x.rstrip() if "t5" in args.model_name else x.rstrip() for x in open(args.input_path).readlines()]
generate_summaries(
examples, args.output_path, args.model_name, batch_size=args.bs, device=args.device, fp16=args.fp16
)
generate_summaries(examples, args.output_path, args.model_name, batch_size=args.bs, device=args.device)
if args.score_path is not None:
output_lns = [x.rstrip() for x in open(args.output_path).readlines()]
reference_lns = [x.rstrip() for x in open(args.reference_path).readlines()]
@@ -23,10 +23,12 @@ logging.basicConfig(level=logging.DEBUG)
logger = logging.getLogger()
FP16_EVER = False
CHEAP_ARGS = {
"theseus_init_copy": False,
"theseus_replace_rate": 0,
"logger": "default",
"num_workers": 2,
"alpha_hid": 0,
"freeze_embeds": True,
"unfreeze_embeds": False,
"enc_only": False,
"tgt_suffix": "",
"resume_from_checkpoint": None,
@@ -72,6 +74,8 @@ CHEAP_ARGS = {
"alpha_loss_encoder": 0.0,
"freeze_encoder": False,
"auto_scale_batch_size": False,
"freeze_decoder": False,
"init_strategy": "bottom",
}
@@ -80,7 +84,6 @@ def _dump_articles(path: Path, articles: list):
f.write("\n".join(articles))
MSG = "T5 is broken at the moment"
T5_TINY = "patrickvonplaten/t5-tiny-random"
@@ -109,8 +112,6 @@ class TestSummarizationDistiller(unittest.TestCase):
freeze_encoder=True,
gpus=2,
sortish_sampler=False,
fp16_opt_level="O1",
fp16=FP16_EVER,
)
self._bart_distiller_cli(updates)
@@ -132,10 +133,6 @@ class TestSummarizationDistiller(unittest.TestCase):
updates = dict(student_encoder_layers=2, student_decoder_layers=1, no_teacher=True,)
self._bart_distiller_cli(updates)
def test_bdc_yes_teacher(self):
updates = dict(student_encoder_layers=2, student_decoder_layers=1,)
self._bart_distiller_cli(updates)
def test_bdc_checkpointing(self):
updates = dict(
student_encoder_layers=2,
@@ -143,6 +140,7 @@ class TestSummarizationDistiller(unittest.TestCase):
num_train_epochs=4,
val_check_interval=0.25,
alpha_hid=2.0,
init_strategy="alternate",
)
model = self._bart_distiller_cli(updates, check_contents=False)
@@ -154,11 +152,31 @@ class TestSummarizationDistiller(unittest.TestCase):
self.assertEqual(len(new_transformer_ckpts), 1)
examples = lmap(str.strip, model.hparams.data_dir.joinpath("test.source").open().readlines())
out_path = tempfile.mktemp()
generate_summaries(examples, out_path, new_transformer_ckpts[0].parent)
generate_summaries(examples, out_path, model_name=str(new_transformer_ckpts[0].parent))
self.assertTrue(Path(out_path).exists())
evaluate_checkpoint(ckpts[0], dest_dir=Path(tempfile.mkdtemp()))
def test_bdc_theseus(self):
updates = dict(
theseus_replace_rate=0.5,
student_encoder_layers=1,
student_decoder_layers=1,
no_teacher=True,
theseus_init_copy=True,
)
self._bart_distiller_cli(updates)
def test_bdc_frozen_theseus(self):
updates = dict(
theseus_replace_rate=0.5,
student_encoder_layers=2,
student_decoder_layers=1,
no_teacher=True,
freeze_encoder=True,
theseus_init_copy=True,
)
self._bart_distiller_cli(updates)
def _bart_distiller_cli(self, updates, check_contents=True):
default_updates = dict(
train_batch_size=1,
+37 -5
View File
@@ -13,6 +13,8 @@ from torch import nn
from torch.utils.data import Dataset, Sampler
from tqdm import tqdm
from transformers import BartTokenizer
def encode_file(
tokenizer,
@@ -30,7 +32,6 @@ def encode_file(
examples = torch.load(cache_path)
assert isinstance(examples, list)
return examples
except Exception:
print(f"failed to load from {cache_path}, retokenizing {data_path}")
data_path = Path(data_path)
@@ -85,7 +86,7 @@ class SummarizationDataset(Dataset):
prefix="",
):
super().__init__()
tok_name = tokenizer.__class__.__name__.lower().rstrip("tokenizer")
tok_name = "T5" if not isinstance(tokenizer, BartTokenizer) else ""
self.source = encode_file(
tokenizer,
os.path.join(data_dir, type_path + ".source"),
@@ -98,6 +99,7 @@ class SummarizationDataset(Dataset):
self.target = encode_file(
tokenizer, tgt_path, max_target_length, overwrite_cache=overwrite_cache, tok_name=tok_name
)
if n_obs is not None:
self.source = self.source[:n_obs]
self.target = self.target[:n_obs]
@@ -212,7 +214,7 @@ def get_git_info():
ROUGE_KEYS = ["rouge1", "rouge2", "rougeL"]
def calculate_rouge(output_lns: List[str], reference_lns: List[str]) -> Dict:
def calculate_rouge(output_lns: List[str], reference_lns: List[str], all_stats=False):
scorer = rouge_scorer.RougeScorer(ROUGE_KEYS, use_stemmer=True)
aggregator = scoring.BootstrapAggregator()
@@ -221,7 +223,11 @@ def calculate_rouge(output_lns: List[str], reference_lns: List[str]) -> Dict:
aggregator.add_scores(scores)
result = aggregator.aggregate()
return {k: v.mid.fmeasure for k, v in result.items()}
if all_stats:
return expanded_rouge_df(result)
else:
return {k: v.mid.fmeasure for k, v in result.items()}
def freeze_params(model: nn.Module):
@@ -241,10 +247,36 @@ def assert_all_frozen(model):
model_grads: List[bool] = list(grad_status(model))
n_require_grad = sum(lmap(int, model_grads))
npars = len(model_grads)
assert not any(model_grads), f"{n_require_grad/npars:.1%} of {npars} weights require grad"
assert not any(model_grads), f"{n_require_grad / npars:.1%} of {npars} weights require grad"
def assert_not_all_frozen(model):
model_grads: List[bool] = list(grad_status(model))
npars = len(model_grads)
assert any(model_grads), f"none of {npars} weights require grad"
def dictify(rouge_obj) -> List:
records = []
for k, rouge_measurement in rouge_obj.items():
if k == "rouge1":
continue
for k1 in ["low", "mid", "high"]:
if k1 != "mid":
continue
v1 = getattr(rouge_measurement, k1)
for k2 in ["precision", "recall", "fmeasure"]:
records.append([k, k1, k2, getattr(v1, k2)])
return records
def expanded_rouge_df(rouge_all):
import pandas as pd
return (
pd.DataFrame(dictify(rouge_all), columns=["metric", "k1", "k2", "val"])
.set_index(["metric", "k2"])["val"]
.unstack("metric")
.rename_axis(None)
)
+40
View File
@@ -0,0 +1,40 @@
exp_name,avg_len,R2_56,R1,R2,RL
dl6_no_teacher,89.17162750217581,,0.4426298740974814,0.21214627660372445,0.30352876002780776
cnn_f12_9_noteach,93.90400348128807,,0.44021766057145817,0.2112684144707048,0.30203109862281485
baseline_140,83.3504,0.2012,0.4405,0.2106,0.3063
dl6_yes_teacher,65.3845953002611,,0.4358592269823179,0.2094872879851778,0.3055267315672064
cnn_12_9_no_teacher,65.63707571801567,,0.4349178365295365,0.20912713468933475,0.30577661504959625
cnn_f12_9,66.25691906005223,,0.4360757701741097,0.20874355812934306,0.30534340603938714
dl6_ckpt,64.7923,0.211,0.4344,0.2087,0.3049
run_6l,64.7923,0.2109,0.4344,0.2086,0.3049
dl6_4ep,66.4003,0.2112,0.4351,0.2082,0.3047
cnn_9_9_yes_teacher,65.75987815491732,,0.43407500508924857,0.2077387062275759,0.3041112175512638
dl6_maxlen_140_gacc2,65.2891,0.2099,0.433,0.2072,0.3038
brewer_cnn_12_6_v2,58.29869451697128,,0.4275596591964453,0.2061398841242274,0.3043299614540603
cnn_enc_only_6_12_brutasse,78.37284595300261,,0.4352447538697767,0.2057028264435689,0.3015751874385552
cnn_no_teacher_f12_3,84.8461270670148,,0.4341908103753934,0.20524584699497528,0.29823737701046504
cnn_6_6_yes_teacher,66.56605744125326,,0.4271799135757925,0.2017168386739932,0.2971576748941705
el6_12_fix,94.3361183637946,,0.4295203853115971,0.20114288940092984,0.29082690956646995
dl3,56.8062,0.2047,0.4197,0.2002,0.2989
dl3_4ep,56.585,0.2057,0.4192,0.1998,0.3
pseudo_dd_140,102.0577,0.1864,0.4254,0.1997,0.2886
dl6_4ep_sasha_ds,60.0613,0.2082,0.4217,0.1993,0.2965
cnn_enc_only_3_12_brutasse,75.5237597911227,,0.4278349430460213,0.19820455870913525,0.2945791235929408
pseudo_6l_140_fix,62.5668,0.2035,0.4197,0.1968,0.2941
staged_cnn_6_6,89.18093994778067,,0.4255412775299003,0.19658951474224035,0.2885276829133131
el9_12,66.26144473455179,,0.42086970025135134,0.1948512909973004,0.29011441089763124
el9_9,65.4414273281114,,0.419385405970708,0.1941718068172067,0.2900372909854251
dl2_4ep,56.013054830287196,,0.4097263081894015,0.1919076904112108,0.29258786794511826
dl3_4ep_sasha_ds,56.9671,0.2007,0.4111,0.1912,0.2914
dl2_4ep_sasha_ds,55.999129677980854,,0.4057876788402908,0.1871524671665716,0.28735289627716404
brewer_cnn_12_1_v4,56.39904264577894,,0.4031071697939543,0.18464802550618464,0.2839190569931007
blarge_12_6_no_teacher,98.89782419495216,,0.4056381963795483,0.1826793955605352,0.2700670904734034
2l_96,56.0208,0.1808,0.395,0.1799,0.2799
2l_140,56.0208,0.1808,0.3951,0.1799,0.2798
2l_56,56.0336,0.1762,0.3895,0.1758,0.2762
2l_eval_140,56.0336,0.1762,0.3897,0.1757,0.2762
el6_12,66.68172323759791,,0.4013734110868821,0.1752702303482339,0.2711560926634892
el6_6,57.5147954743255,,0.38798696359999946,0.1672810503629001,0.2648421408851355
dl1,56.0179,0.1264,0.3411,0.1285,0.226
last_layer,55.6623,0.0596,0.2654,0.0616,0.1764
cnn_enc_only_9_12_brutasse,122.32297650130549,,0.12110269669223532,0.011040620587715645,0.09900079897379864
1 exp_name avg_len R2_56 R1 R2 RL
2 dl6_no_teacher 89.17162750217581 0.4426298740974814 0.21214627660372445 0.30352876002780776
3 cnn_f12_9_noteach 93.90400348128807 0.44021766057145817 0.2112684144707048 0.30203109862281485
4 baseline_140 83.3504 0.2012 0.4405 0.2106 0.3063
5 dl6_yes_teacher 65.3845953002611 0.4358592269823179 0.2094872879851778 0.3055267315672064
6 cnn_12_9_no_teacher 65.63707571801567 0.4349178365295365 0.20912713468933475 0.30577661504959625
7 cnn_f12_9 66.25691906005223 0.4360757701741097 0.20874355812934306 0.30534340603938714
8 dl6_ckpt 64.7923 0.211 0.4344 0.2087 0.3049
9 run_6l 64.7923 0.2109 0.4344 0.2086 0.3049
10 dl6_4ep 66.4003 0.2112 0.4351 0.2082 0.3047
11 cnn_9_9_yes_teacher 65.75987815491732 0.43407500508924857 0.2077387062275759 0.3041112175512638
12 dl6_maxlen_140_gacc2 65.2891 0.2099 0.433 0.2072 0.3038
13 brewer_cnn_12_6_v2 58.29869451697128 0.4275596591964453 0.2061398841242274 0.3043299614540603
14 cnn_enc_only_6_12_brutasse 78.37284595300261 0.4352447538697767 0.2057028264435689 0.3015751874385552
15 cnn_no_teacher_f12_3 84.8461270670148 0.4341908103753934 0.20524584699497528 0.29823737701046504
16 cnn_6_6_yes_teacher 66.56605744125326 0.4271799135757925 0.2017168386739932 0.2971576748941705
17 el6_12_fix 94.3361183637946 0.4295203853115971 0.20114288940092984 0.29082690956646995
18 dl3 56.8062 0.2047 0.4197 0.2002 0.2989
19 dl3_4ep 56.585 0.2057 0.4192 0.1998 0.3
20 pseudo_dd_140 102.0577 0.1864 0.4254 0.1997 0.2886
21 dl6_4ep_sasha_ds 60.0613 0.2082 0.4217 0.1993 0.2965
22 cnn_enc_only_3_12_brutasse 75.5237597911227 0.4278349430460213 0.19820455870913525 0.2945791235929408
23 pseudo_6l_140_fix 62.5668 0.2035 0.4197 0.1968 0.2941
24 staged_cnn_6_6 89.18093994778067 0.4255412775299003 0.19658951474224035 0.2885276829133131
25 el9_12 66.26144473455179 0.42086970025135134 0.1948512909973004 0.29011441089763124
26 el9_9 65.4414273281114 0.419385405970708 0.1941718068172067 0.2900372909854251
27 dl2_4ep 56.013054830287196 0.4097263081894015 0.1919076904112108 0.29258786794511826
28 dl3_4ep_sasha_ds 56.9671 0.2007 0.4111 0.1912 0.2914
29 dl2_4ep_sasha_ds 55.999129677980854 0.4057876788402908 0.1871524671665716 0.28735289627716404
30 brewer_cnn_12_1_v4 56.39904264577894 0.4031071697939543 0.18464802550618464 0.2839190569931007
31 blarge_12_6_no_teacher 98.89782419495216 0.4056381963795483 0.1826793955605352 0.2700670904734034
32 2l_96 56.0208 0.1808 0.395 0.1799 0.2799
33 2l_140 56.0208 0.1808 0.3951 0.1799 0.2798
34 2l_56 56.0336 0.1762 0.3895 0.1758 0.2762
35 2l_eval_140 56.0336 0.1762 0.3897 0.1757 0.2762
36 el6_12 66.68172323759791 0.4013734110868821 0.1752702303482339 0.2711560926634892
37 el6_6 57.5147954743255 0.38798696359999946 0.1672810503629001 0.2648421408851355
38 dl1 56.0179 0.1264 0.3411 0.1285 0.226
39 last_layer 55.6623 0.0596 0.2654 0.0616 0.1764
40 cnn_enc_only_9_12_brutasse 122.32297650130549 0.12110269669223532 0.011040620587715645 0.09900079897379864
+30
View File
@@ -0,0 +1,30 @@
exp_name,avg_len,R1,R2,RL
brewer_xsum_12_6,27.251213270978557,0.45257522039021847,0.2195833165680984,0.3679937489447971
xsum_baseline,29.06997264625431,0.45234885716476203,0.21848808526257507,0.3650268517677788
xsum_f12_9,29.35074561016501,0.451213614023561,0.2163378899615816,0.363539265208816
brewer_xsum_9_6,27.35771640342363,0.4492537548158509,0.2162358630026469,0.3641728953722555
xsum_no_teacher_12_6,27.75319862348893,0.4486508407826425,0.2149919319149121,0.36419730252620985
xsum_clean_baseline,29.314656313420983,0.4477747607857255,0.21439354376130093,0.3593759295843251
brewer_xsum_12_3,24.899232330362658,0.44368280928513437,0.21344371120165062,0.3630886297392266
xsum_f12_9_noteach,29.46845495455749,0.4472052566240151,0.21274301655819566,0.3593763343512352
xsum_12_9_yes_teacher_mlm_low,29.737845230742078,0.4465496602405639,0.2123133092772882,0.35962822466043964
brewer_xsum_9_9_v8.2,27.76872849201447,0.4443257538168877,0.21127297599326367,0.3582479791521699
xsum_theseus_12_6_copy,26.884320127062562,0.4424918692803468,0.21067481121631593,0.3583553330701445
brewer_xsum_9_6_gaccum,27.142063001853,0.4425169154001685,0.2085901108525016,0.3575200377038006
xsum_dl6,27.928086120180005,0.4428406843456235,0.20797826365160146,0.3570522063056493
brewer_xsum_6_6_v3,27.265684284831906,0.43971173638675576,0.20748396024553475,0.3548038868729512
xsum_no_teacher_f12_6,27.745874878672904,0.4410138016184357,0.20718399111770808,0.3554781824669484
xsum_9_9_encloss10,28.9296744021883,0.4375314579000953,0.2042900182424192,0.3507282669577781
xsum_enc_only_9_12_v8,27.90267360804729,0.4317712176227767,0.1975871213673145,0.3445507930002034
xsum_no_teacher_f12_3,25.100238242301245,0.4212459339899797,0.19363224077581093,0.3425534541928968
xsum_enc_only_6_12_brutasse,28.53339804111885,0.4268136827870216,0.19356708317121352,0.3391381862315658
xsum_9_12_noteacher,29.28192005647225,0.4197414451047208,0.1873359194795431,0.33264079827956505
eval_xsum_9_9_no_teacher,28.378540545310155,0.4154520215812441,0.1838431267358389,0.3296500161194308
xsum_teacher_f12_3,22.863937174622784,0.4099647311777451,0.1838389871897501,0.3350523158599553
xsum_9_9_no_teacher,28.378540545310155,0.4155138730166103,0.1837896883494328,0.3296699124599892
xsum_6_6_yes_teacher,27.355245742521838,0.4156289447797342,0.1833851442061078,0.3312280244052153
brewer_xsum_3_3_v2,24.55272213888644,0.4066228194707484,0.1820601762610753,0.3289002584850459
xsum_theseus_6_6_v2,27.172240360010587,0.3843975566079777,0.16290567759393396,0.3068524775579614
brewer_xsum_6_6_fix2,26.134651019147622,0.38678850503910334,0.15887838687468492,0.3076899937354824
staged_xsum_9_9,27.46360187064325,0.3735212492621201,0.14991821061038266,0.2934425279687253
brewer_xsum_6_6_newcode,59.600811788582014,0.24447308274539264,0.06887251729221451,0.17840798937780264
1 exp_name avg_len R1 R2 RL
2 brewer_xsum_12_6 27.251213270978557 0.45257522039021847 0.2195833165680984 0.3679937489447971
3 xsum_baseline 29.06997264625431 0.45234885716476203 0.21848808526257507 0.3650268517677788
4 xsum_f12_9 29.35074561016501 0.451213614023561 0.2163378899615816 0.363539265208816
5 brewer_xsum_9_6 27.35771640342363 0.4492537548158509 0.2162358630026469 0.3641728953722555
6 xsum_no_teacher_12_6 27.75319862348893 0.4486508407826425 0.2149919319149121 0.36419730252620985
7 xsum_clean_baseline 29.314656313420983 0.4477747607857255 0.21439354376130093 0.3593759295843251
8 brewer_xsum_12_3 24.899232330362658 0.44368280928513437 0.21344371120165062 0.3630886297392266
9 xsum_f12_9_noteach 29.46845495455749 0.4472052566240151 0.21274301655819566 0.3593763343512352
10 xsum_12_9_yes_teacher_mlm_low 29.737845230742078 0.4465496602405639 0.2123133092772882 0.35962822466043964
11 brewer_xsum_9_9_v8.2 27.76872849201447 0.4443257538168877 0.21127297599326367 0.3582479791521699
12 xsum_theseus_12_6_copy 26.884320127062562 0.4424918692803468 0.21067481121631593 0.3583553330701445
13 brewer_xsum_9_6_gaccum 27.142063001853 0.4425169154001685 0.2085901108525016 0.3575200377038006
14 xsum_dl6 27.928086120180005 0.4428406843456235 0.20797826365160146 0.3570522063056493
15 brewer_xsum_6_6_v3 27.265684284831906 0.43971173638675576 0.20748396024553475 0.3548038868729512
16 xsum_no_teacher_f12_6 27.745874878672904 0.4410138016184357 0.20718399111770808 0.3554781824669484
17 xsum_9_9_encloss10 28.9296744021883 0.4375314579000953 0.2042900182424192 0.3507282669577781
18 xsum_enc_only_9_12_v8 27.90267360804729 0.4317712176227767 0.1975871213673145 0.3445507930002034
19 xsum_no_teacher_f12_3 25.100238242301245 0.4212459339899797 0.19363224077581093 0.3425534541928968
20 xsum_enc_only_6_12_brutasse 28.53339804111885 0.4268136827870216 0.19356708317121352 0.3391381862315658
21 xsum_9_12_noteacher 29.28192005647225 0.4197414451047208 0.1873359194795431 0.33264079827956505
22 eval_xsum_9_9_no_teacher 28.378540545310155 0.4154520215812441 0.1838431267358389 0.3296500161194308
23 xsum_teacher_f12_3 22.863937174622784 0.4099647311777451 0.1838389871897501 0.3350523158599553
24 xsum_9_9_no_teacher 28.378540545310155 0.4155138730166103 0.1837896883494328 0.3296699124599892
25 xsum_6_6_yes_teacher 27.355245742521838 0.4156289447797342 0.1833851442061078 0.3312280244052153
26 brewer_xsum_3_3_v2 24.55272213888644 0.4066228194707484 0.1820601762610753 0.3289002584850459
27 xsum_theseus_6_6_v2 27.172240360010587 0.3843975566079777 0.16290567759393396 0.3068524775579614
28 brewer_xsum_6_6_fix2 26.134651019147622 0.38678850503910334 0.15887838687468492 0.3076899937354824
29 staged_xsum_9_9 27.46360187064325 0.3735212492621201 0.14991821061038266 0.2934425279687253
30 brewer_xsum_6_6_newcode 59.600811788582014 0.24447308274539264 0.06887251729221451 0.17840798937780264
+8 -1
View File
@@ -69,6 +69,9 @@ class BartConfig(PretrainedConfig):
normalize_embedding=True,
static_position_embeddings=False,
add_bias_logits=False,
student_decoder_layers=None,
student_encoder_layers=None,
replacing_rate=0,
**common_kwargs
):
r"""
@@ -87,6 +90,7 @@ class BartConfig(PretrainedConfig):
is_encoder_decoder=is_encoder_decoder,
**common_kwargs,
)
self.replacing_rate = replacing_rate
self.vocab_size = vocab_size
self.d_model = d_model # encoder_embed_dim and decoder_embed_dim
self.encoder_ffn_dim = encoder_ffn_dim
@@ -119,9 +123,12 @@ class BartConfig(PretrainedConfig):
# Classifier stuff
self.classif_dropout = classifier_dropout
# pos embedding offset
self.extra_pos_embeddings = self.pad_token_id + 1
# Theseus params
self.student_encoder_layers = student_encoder_layers
self.student_decoder_layers = student_decoder_layers
@property
def num_attention_heads(self) -> int:
return self.encoder_attention_heads
+70 -9
View File
@@ -23,6 +23,7 @@ import numpy as np
import torch
import torch.nn.functional as F
from torch import Tensor, nn
from torch.distributions.bernoulli import Bernoulli
from torch.nn import CrossEntropyLoss
from .activations import ACT2FN
@@ -44,6 +45,23 @@ BART_PRETRAINED_MODEL_ARCHIVE_LIST = [
]
def get_layers_to_copy(n_to_get, tot):
all_layers = list(range(tot))
if tot == 12: # Alternating for special cases
layers_to_copy = { # maps # layers in student -> which teacher layers to copy
6: [0, 2, 4, 7, 9, 11],
1: [11],
3: [0, 6, 11],
2: [0, 11],
4: [0, 4, 8, 11],
9: [0, 1, 2, 4, 5, 7, 9, 10, 11],
12: all_layers,
}
return layers_to_copy[n_to_get]
else:
return all_layers[:n_to_get]
BART_START_DOCSTRING = r"""
This model is a PyTorch `torch.nn.Module <https://pytorch.org/docs/stable/nn.html#torch.nn.Module>`_ sub-class. Use it as a regular PyTorch Module and
@@ -62,6 +80,7 @@ BART_GENERATION_EXAMPLE = r"""
# see ``examples/summarization/bart/run_eval.py`` for a longer example
model = BartForConditionalGeneration.from_pretrained('facebook/bart-large-cnn')
tokenizer = BartTokenizer.from_pretrained('facebook/bart-large-cnn')
ARTICLE_TO_SUMMARIZE = "My friends are cool but they eat too many carbs."
inputs = tokenizer.batch_encode_plus([ARTICLE_TO_SUMMARIZE], max_length=1024, return_tensors='pt')
# Generate Summary
@@ -235,7 +254,37 @@ class EncoderLayer(nn.Module):
return x, attn_weights
class BartEncoder(nn.Module):
class TheseusMixin:
compress_ratio = 2
def set_replacing_rate(self, replacing_rate):
if not 0 < replacing_rate <= 1:
raise Exception("Replace rate must be in the range (0, 1]!")
self.bernoulli = Bernoulli(torch.tensor([replacing_rate]))
def determine_inference_layers(self) -> nn.ModuleList:
if self.scc_layers is None:
return self.layers
if self.training:
inference_layers = []
for i in range(len(self.scc_layers)):
if self.bernoulli.sample() == 1: # REPLACE
inference_layers.append(self.scc_layers[i])
else: # KEEP the original
for offset in range(self.compress_ratio):
inference_layers.append(self.layers[i * self.compress_ratio + offset])
else: # inference with compressed model
inference_layers = self.scc_layers
return inference_layers
def init_successor_layers(self, replacing_rate):
if replacing_rate > 0 and replacing_rate < 1:
self.set_replacing_rate(replacing_rate)
class BartEncoder(nn.Module, TheseusMixin):
"""
Transformer encoder consisting of *config.encoder_layers* self attention layers. Each layer
is a :class:`EncoderLayer`.
@@ -265,9 +314,15 @@ class BartEncoder(nn.Module):
config.max_position_embeddings, embed_dim, self.padding_idx, config.extra_pos_embeddings,
)
self.layers = nn.ModuleList([EncoderLayer(config) for _ in range(config.encoder_layers)])
self.layernorm_embedding = LayerNorm(embed_dim) if config.normalize_embedding else nn.Identity()
# mbart has one extra layer_norm
self.layer_norm = LayerNorm(config.d_model) if config.normalize_before else None
self.scc_layers = None
if config.student_encoder_layers is not None and config.student_encoder_layers < config.encoder_layers:
self.scc_layers = nn.ModuleList([EncoderLayer(config) for _ in range(config.student_encoder_layers)])
self.compress_ratio = len(self.scc_layers) // len(self.layers)
self.init_successor_layers(config.replacing_rate)
def forward(self, input_ids, attention_mask=None, output_attentions=False, output_hidden_states=False):
"""
@@ -294,12 +349,13 @@ class BartEncoder(nn.Module):
x = inputs_embeds + embed_pos
x = self.layernorm_embedding(x)
x = F.dropout(x, p=self.dropout, training=self.training)
inference_layers = self.determine_inference_layers()
# B x T x C -> T x B x C
x = x.transpose(0, 1)
encoder_states, all_attentions = [], []
for encoder_layer in self.layers:
for encoder_layer in inference_layers:
if output_hidden_states:
encoder_states.append(x)
# add LayerDrop (see https://arxiv.org/abs/1909.11556 for description)
@@ -413,7 +469,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, TheseusMixin):
"""
Transformer decoder consisting of *config.decoder_layers* layers. Each layer
is a :class:`DecoderLayer`.
@@ -443,6 +499,11 @@ class BartDecoder(nn.Module):
) # type: List[DecoderLayer]
self.layernorm_embedding = LayerNorm(config.d_model) if config.normalize_embedding else nn.Identity()
self.layer_norm = LayerNorm(config.d_model) if config.add_final_layer_norm else None
self.scc_layers = None
if config.student_decoder_layers is not None and config.student_decoder_layers < config.decoder_layers:
self.scc_layers = nn.ModuleList([DecoderLayer(config) for _ in range(config.student_decoder_layers)])
self.compress_ratio = len(self.scc_layers) // len(self.layers)
self.init_successor_layers(config.replacing_rate)
def forward(
self,
@@ -500,7 +561,9 @@ class BartDecoder(nn.Module):
all_hidden_states = ()
all_self_attns = ()
next_decoder_cache = []
for idx, decoder_layer in enumerate(self.layers):
inference_layers = self.determine_inference_layers()
for idx, decoder_layer in enumerate(inference_layers):
# add LayerDrop (see https://arxiv.org/abs/1909.11556 for description)
if output_hidden_states:
all_hidden_states += (x,)
@@ -509,7 +572,7 @@ class BartDecoder(nn.Module):
continue
layer_state = decoder_cached_states[idx] if decoder_cached_states is not None else None
# DecoderLayer.forward()
x, layer_self_attn, layer_past = decoder_layer(
x,
encoder_hidden_states,
@@ -797,7 +860,6 @@ def _get_shape(t):
class BartModel(PretrainedBartModel):
def __init__(self, config: BartConfig):
super().__init__(config)
padding_idx, vocab_size = config.pad_token_id, config.vocab_size
self.shared = nn.Embedding(vocab_size, config.d_model, padding_idx)
@@ -976,8 +1038,8 @@ class BartForConditionalGeneration(PretrainedBartModel):
decoder_attention_mask=decoder_attention_mask,
decoder_cached_states=decoder_cached_states,
use_cache=use_cache,
output_attentions=output_attentions,
output_hidden_states=output_hidden_states,
output_attentions=output_attentions,
)
lm_logits = F.linear(outputs[0], self.model.shared.weight, bias=self.final_logits_bias)
outputs = (lm_logits,) + outputs[1:] # Add cache, hidden states and attention if they are here
@@ -991,7 +1053,6 @@ class BartForConditionalGeneration(PretrainedBartModel):
def prepare_inputs_for_generation(self, decoder_input_ids, past, attention_mask, use_cache, **kwargs):
assert past is not None, "past has to be defined for encoder_outputs"
encoder_outputs, decoder_cached_states = past
return {
"input_ids": None, # encoder_outputs is defined. input_ids not needed
@@ -1112,8 +1173,8 @@ class BartForSequenceClassification(PretrainedBartModel):
decoder_input_ids=decoder_input_ids,
decoder_attention_mask=decoder_attention_mask,
encoder_outputs=encoder_outputs,
output_attentions=output_attentions,
output_hidden_states=output_hidden_states,
output_attentions=output_attentions,
)
x = outputs[0] # last hidden state
eos_mask = input_ids.eq(self.config.eos_token_id)
+19
View File
@@ -0,0 +1,19 @@
,length_penalty,max_length,min_length,num_beams,rouge1,rouge2,rougeL
beam_search_xsum_base_500/generations_0.txt,0.5,62.0,11.0,2.0,0.4405564669595006,0.1983002888180257,0.3437812408347012
beam_search_xsum_base_500/generations_1.txt,0.5,62.0,11.0,4.0,0.4448474889733566,0.20778679115689908,0.35315548966820265
beam_search_xsum_base_500/generations_2.txt,0.5,62.0,11.0,6.0,0.44523636974022285,0.21033071404850712,0.3580167024803249
beam_search_xsum_base_500/generations_3.txt,0.5,62.0,11.0,8.0,0.4468688924754079,0.21170992719882548,0.35932945274070693
beam_search_xsum_base_500/generations_4.txt,0.5,62.0,11.0,10.0,0.44549478467300774,0.21288210214860998,0.36011253722710324
beam_search_xsum_base_500/generations_5.txt,0.5,62.0,11.0,20.0,0.4459196355164726,0.21553443653166177,0.36066326922139236
beam_search_xsum_base_500/generations_6.txt,1.0,62.0,11.0,2.0,0.44053092513361236,0.1993972468658327,0.3444322842619076
beam_search_xsum_base_500/generations_7.txt,1.0,62.0,11.0,4.0,0.4417046134948439,0.20676749488131385,0.3499947350751297
beam_search_xsum_base_500/generations_8.txt,1.0,62.0,11.0,6.0,0.44452835581665295,0.20896315664629483,0.3549064385172764
beam_search_xsum_base_500/generations_9.txt,1.0,62.0,11.0,8.0,0.44358288615105523,0.2093555938220097,0.3549769613045229
beam_search_xsum_base_500/generations_10.txt,1.0,62.0,11.0,10.0,0.44330114244232066,0.2119742308397891,0.3570212060039547
beam_search_xsum_base_500/generations_11.txt,1.0,62.0,11.0,20.0,0.44368891403518096,0.21275693156665768,0.35774144153552306
beam_search_xsum_base_500/generations_12.txt,2.0,62.0,11.0,2.0,0.4351172939026171,0.1946192131538515,0.338306269234114
beam_search_xsum_base_500/generations_13.txt,2.0,62.0,11.0,4.0,0.4355928169879645,0.20028529354739374,0.3409900672696907
beam_search_xsum_base_500/generations_14.txt,2.0,62.0,11.0,6.0,0.43613272840789125,0.20204391577628494,0.34330389423389684
beam_search_xsum_base_500/generations_15.txt,2.0,62.0,11.0,8.0,0.43189181564990875,0.1985819991945317,0.33890611664212345
beam_search_xsum_base_500/generations_16.txt,2.0,62.0,11.0,10.0,0.4312709248247222,0.19786145957678045,0.340228104937922
beam_search_xsum_base_500/generations_17.txt,2.0,62.0,11.0,20.0,0.431958235502211,0.1997066827574804,0.3410722313510985
1 length_penalty max_length min_length num_beams rouge1 rouge2 rougeL
2 beam_search_xsum_base_500/generations_0.txt 0.5 62.0 11.0 2.0 0.4405564669595006 0.1983002888180257 0.3437812408347012
3 beam_search_xsum_base_500/generations_1.txt 0.5 62.0 11.0 4.0 0.4448474889733566 0.20778679115689908 0.35315548966820265
4 beam_search_xsum_base_500/generations_2.txt 0.5 62.0 11.0 6.0 0.44523636974022285 0.21033071404850712 0.3580167024803249
5 beam_search_xsum_base_500/generations_3.txt 0.5 62.0 11.0 8.0 0.4468688924754079 0.21170992719882548 0.35932945274070693
6 beam_search_xsum_base_500/generations_4.txt 0.5 62.0 11.0 10.0 0.44549478467300774 0.21288210214860998 0.36011253722710324
7 beam_search_xsum_base_500/generations_5.txt 0.5 62.0 11.0 20.0 0.4459196355164726 0.21553443653166177 0.36066326922139236
8 beam_search_xsum_base_500/generations_6.txt 1.0 62.0 11.0 2.0 0.44053092513361236 0.1993972468658327 0.3444322842619076
9 beam_search_xsum_base_500/generations_7.txt 1.0 62.0 11.0 4.0 0.4417046134948439 0.20676749488131385 0.3499947350751297
10 beam_search_xsum_base_500/generations_8.txt 1.0 62.0 11.0 6.0 0.44452835581665295 0.20896315664629483 0.3549064385172764
11 beam_search_xsum_base_500/generations_9.txt 1.0 62.0 11.0 8.0 0.44358288615105523 0.2093555938220097 0.3549769613045229
12 beam_search_xsum_base_500/generations_10.txt 1.0 62.0 11.0 10.0 0.44330114244232066 0.2119742308397891 0.3570212060039547
13 beam_search_xsum_base_500/generations_11.txt 1.0 62.0 11.0 20.0 0.44368891403518096 0.21275693156665768 0.35774144153552306
14 beam_search_xsum_base_500/generations_12.txt 2.0 62.0 11.0 2.0 0.4351172939026171 0.1946192131538515 0.338306269234114
15 beam_search_xsum_base_500/generations_13.txt 2.0 62.0 11.0 4.0 0.4355928169879645 0.20028529354739374 0.3409900672696907
16 beam_search_xsum_base_500/generations_14.txt 2.0 62.0 11.0 6.0 0.43613272840789125 0.20204391577628494 0.34330389423389684
17 beam_search_xsum_base_500/generations_15.txt 2.0 62.0 11.0 8.0 0.43189181564990875 0.1985819991945317 0.33890611664212345
18 beam_search_xsum_base_500/generations_16.txt 2.0 62.0 11.0 10.0 0.4312709248247222 0.19786145957678045 0.340228104937922
19 beam_search_xsum_base_500/generations_17.txt 2.0 62.0 11.0 20.0 0.431958235502211 0.1997066827574804 0.3410722313510985
+19
View File
@@ -0,0 +1,19 @@
,length_penalty,max_length,min_length,num_beams,time,rouge1,rouge2,rougeL
beam_search_xsum_d6_500/generations_5.txt,1.0,62.0,11.0,20.0,236.1429727077484,0.4305654553611711,0.19600039925703722,0.34201476194014246
beam_search_xsum_d6_500/generations_4.txt,1.0,62.0,11.0,10.0,135.6617534160614,0.42671786797903843,0.1929509089620034,0.3357873765545949
beam_search_xsum_d6_500/generations_3.txt,1.0,62.0,11.0,8.0,118.18162441253662,0.4244918905164582,0.19240953848063766,0.33592718733437477
beam_search_xsum_d6_500/generations_2.txt,1.0,62.0,11.0,6.0,99.75846719741821,0.4241986674718551,0.1904178221737402,0.33554345459177803
beam_search_xsum_d6_500/generations_1.txt,1.0,62.0,11.0,4.0,82.19014501571655,0.42552162916374947,0.1900029083958625,0.3348035451647859
beam_search_xsum_d6_500/generations_0.txt,1.0,62.0,11.0,2.0,73.42616081237793,0.42185608960262033,0.188239694825568,0.3380120701696464
beam_search_xsum_d6_500/generations_10.txt,2.0,62.0,11.0,10.0,134.8128764629364,0.41995279044460687,0.1873177328806112,0.32606274473224905
beam_search_xsum_d6_500/generations_16.txt,3.0,62.0,11.0,10.0,135.02988576889038,0.41958442037714705,0.1862058575783302,0.3248733090856767
beam_search_xsum_d6_500/generations_11.txt,2.0,62.0,11.0,20.0,235.85728216171265,0.4218916450932465,0.18532554772332022,0.3277383511244094
beam_search_xsum_d6_500/generations_8.txt,2.0,62.0,11.0,6.0,97.2721335887909,0.41963315273987645,0.18522376636363672,0.3271210646772599
beam_search_xsum_d6_500/generations_9.txt,2.0,62.0,11.0,8.0,115.15277361869812,0.417398394462531,0.1845339882603932,0.3254173997315589
beam_search_xsum_d6_500/generations_15.txt,3.0,62.0,11.0,8.0,115.81080627441406,0.41703499851068504,0.18357496523072026,0.3242598639398694
beam_search_xsum_d6_500/generations_6.txt,2.0,62.0,11.0,2.0,68.16163468360901,0.41780417025178973,0.18351703902053906,0.3313922817377364
beam_search_xsum_d6_500/generations_14.txt,3.0,62.0,11.0,6.0,97.10926246643066,0.41681309969268543,0.18350802133208385,0.3248941386957547
beam_search_xsum_d6_500/generations_7.txt,2.0,62.0,11.0,4.0,79.3750171661377,0.4192827070533588,0.18286781380405798,0.32295416688769185
beam_search_xsum_d6_500/generations_12.txt,3.0,62.0,11.0,2.0,67.85268831253052,0.4173974464414504,0.18261800246898374,0.33032670326211594
beam_search_xsum_d6_500/generations_13.txt,3.0,62.0,11.0,4.0,78.7077898979187,0.41858355068804654,0.18122511913031086,0.32144676793629146
beam_search_xsum_d6_500/generations_17.txt,3.0,62.0,11.0,20.0,236.6838607788086,0.25952276985447975,0.100359902189957,0.19233054575018543
1 length_penalty max_length min_length num_beams time rouge1 rouge2 rougeL
2 beam_search_xsum_d6_500/generations_5.txt 1.0 62.0 11.0 20.0 236.1429727077484 0.4305654553611711 0.19600039925703722 0.34201476194014246
3 beam_search_xsum_d6_500/generations_4.txt 1.0 62.0 11.0 10.0 135.6617534160614 0.42671786797903843 0.1929509089620034 0.3357873765545949
4 beam_search_xsum_d6_500/generations_3.txt 1.0 62.0 11.0 8.0 118.18162441253662 0.4244918905164582 0.19240953848063766 0.33592718733437477
5 beam_search_xsum_d6_500/generations_2.txt 1.0 62.0 11.0 6.0 99.75846719741821 0.4241986674718551 0.1904178221737402 0.33554345459177803
6 beam_search_xsum_d6_500/generations_1.txt 1.0 62.0 11.0 4.0 82.19014501571655 0.42552162916374947 0.1900029083958625 0.3348035451647859
7 beam_search_xsum_d6_500/generations_0.txt 1.0 62.0 11.0 2.0 73.42616081237793 0.42185608960262033 0.188239694825568 0.3380120701696464
8 beam_search_xsum_d6_500/generations_10.txt 2.0 62.0 11.0 10.0 134.8128764629364 0.41995279044460687 0.1873177328806112 0.32606274473224905
9 beam_search_xsum_d6_500/generations_16.txt 3.0 62.0 11.0 10.0 135.02988576889038 0.41958442037714705 0.1862058575783302 0.3248733090856767
10 beam_search_xsum_d6_500/generations_11.txt 2.0 62.0 11.0 20.0 235.85728216171265 0.4218916450932465 0.18532554772332022 0.3277383511244094
11 beam_search_xsum_d6_500/generations_8.txt 2.0 62.0 11.0 6.0 97.2721335887909 0.41963315273987645 0.18522376636363672 0.3271210646772599
12 beam_search_xsum_d6_500/generations_9.txt 2.0 62.0 11.0 8.0 115.15277361869812 0.417398394462531 0.1845339882603932 0.3254173997315589
13 beam_search_xsum_d6_500/generations_15.txt 3.0 62.0 11.0 8.0 115.81080627441406 0.41703499851068504 0.18357496523072026 0.3242598639398694
14 beam_search_xsum_d6_500/generations_6.txt 2.0 62.0 11.0 2.0 68.16163468360901 0.41780417025178973 0.18351703902053906 0.3313922817377364
15 beam_search_xsum_d6_500/generations_14.txt 3.0 62.0 11.0 6.0 97.10926246643066 0.41681309969268543 0.18350802133208385 0.3248941386957547
16 beam_search_xsum_d6_500/generations_7.txt 2.0 62.0 11.0 4.0 79.3750171661377 0.4192827070533588 0.18286781380405798 0.32295416688769185
17 beam_search_xsum_d6_500/generations_12.txt 3.0 62.0 11.0 2.0 67.85268831253052 0.4173974464414504 0.18261800246898374 0.33032670326211594
18 beam_search_xsum_d6_500/generations_13.txt 3.0 62.0 11.0 4.0 78.7077898979187 0.41858355068804654 0.18122511913031086 0.32144676793629146
19 beam_search_xsum_d6_500/generations_17.txt 3.0 62.0 11.0 20.0 236.6838607788086 0.25952276985447975 0.100359902189957 0.19233054575018543