Compare commits

...
Author SHA1 Message Date
jetrunner 20dd34f887 Switch prd & scc for LayerDrop & mixout 2020-09-17 22:38:10 +08:00
jetrunner 181b4c28a8 md 2020-09-17 22:38:10 +08:00
jetrunner 3cf1064932 TF tests 2020-09-17 22:38:10 +08:00
jetrunner d0c53ad1f6 readme 2020-09-17 22:38:10 +08:00
jetrunner 33c569f0f1 change the behavior of forward 2020-09-17 22:38:10 +08:00
jetrunner d66af94fa4 make style 2020-09-17 22:38:10 +08:00
jetrunner 4cfafd6bdf Add unit tests 2020-09-17 22:38:10 +08:00
jetrunner b73d1e6e7f flake 8 2020-09-17 22:38:10 +08:00
jetrunner fb8dabc2dd make style 2020-09-17 22:38:10 +08:00
jetrunner 934139678b Add Theseus Compression 2020-09-17 22:38:10 +08:00
8 changed files with 290 additions and 0 deletions
+28
View File
@@ -0,0 +1,28 @@
# transformers.theseus
![BERT of Theseus](https://github.com/JetRunner/BERT-of-Theseus/blob/master/bert-of-theseus.png?raw=true)
`transformers.theseus` is an implementation of [Theseus Compression](https://arxiv.org/abs/2002.02925).
Theseus compression exploits module replacing to compress a large model to a small one.
We implement [LayerDrop](https://arxiv.org/abs/1909.11556) and [Mixout](https://arxiv.org/abs/1909.11299) with the framework as well.
## Run BERT-of-Theseus
TBA.
## Run LayerDrop
TBA.
## Run Mixout
TBA.
## Citation
Please consider citing this paper for the `theseus` framework and BERT-of-Theseus:
```bibtex
@misc{xu2020bertoftheseus,
title={BERT-of-Theseus: Compressing BERT by Progressive Module Replacing},
author={Canwen Xu and Wangchunshu Zhou and Tao Ge and Furu Wei and Ming Zhou},
year={2020},
eprint={2002.02925},
archivePrefix={arXiv},
primaryClass={cs.CL}
}
```
+6
View File
@@ -0,0 +1,6 @@
# flake8: noqa
from .layerdrop_list import LayerDropList
from .mixout_list import MixoutList
from .theseus_list import TheseusList
from .theseus_module import TheseusModule
@@ -0,0 +1,15 @@
from .theseus_list import TheseusList
from .theseus_module import TheseusModule
class LayerDropList(TheseusList):
"""
Implementation of Layer Drop (https://arxiv.org/abs/1909.11556).
"""
@classmethod
def from_module_list(cls, module_list, replacing_rate):
list_to_return = cls()
for module in module_list:
list_to_return.append(TheseusModule(successor=module, replacing_rate=replacing_rate))
return list_to_return
+29
View File
@@ -0,0 +1,29 @@
from copy import deepcopy
from .theseus_list import TheseusList
from .theseus_module import TheseusModule
class MixoutList(TheseusList):
"""
Implementation of Mixout (https://arxiv.org/abs/1909.11299).
"""
@classmethod
def from_module_list(cls, module_list, replacing_rate, freeze_predecessor=True):
"""
:param module_list:
:param replacing_rate:
:param freeze_predecessor: whether to freeze the original pretraining weights.
:return:
"""
list_to_return = cls()
for module in module_list:
predecessor = deepcopy(module)
if freeze_predecessor:
for param in predecessor.parameters():
param.requires_grad = False
list_to_return.append(
TheseusModule(predecessor=predecessor, successor=module, replacing_rate=replacing_rate)
)
return list_to_return
@@ -0,0 +1,2 @@
class NoSuccessorError(Exception):
pass
+47
View File
@@ -0,0 +1,47 @@
import torch
from .theseus_module import TheseusModule
def _unpack_module(packed_module):
list_to_return = torch.nn.ModuleList()
if isinstance(packed_module, (list, tuple, torch.nn.ModuleList)):
for submodule in packed_module:
list_to_return.append(submodule)
elif isinstance(packed_module, torch.nn.Module):
list_to_return.append(packed_module)
return list_to_return
class TheseusList(torch.nn.ModuleList):
"""
TheseusList is a ModuleList that implements methods for Theseus Compression.
"""
def set_replacing_rate(self, replacing_rate):
for module in self:
if isinstance(module, TheseusModule):
module.set_replacing_rate(replacing_rate)
def sample_and_pass(self) -> torch.nn.ModuleList:
list_to_return = torch.nn.ModuleList()
for module in self:
if isinstance(module, TheseusModule):
list_to_return += _unpack_module(module.sample_and_pass())
else:
list_to_return += _unpack_module(module)
return list_to_return
def get_successors(self) -> torch.nn.ModuleList:
list_to_return = torch.nn.ModuleList()
for module in self:
if isinstance(module, TheseusModule) and module.successor:
list_to_return += _unpack_module(module.successor)
return list_to_return
def get_predecessors(self) -> torch.nn.ModuleList:
list_to_return = torch.nn.ModuleList()
for module in self:
if isinstance(module, TheseusModule):
list_to_return += _unpack_module(module.predecessor)
return list_to_return
@@ -0,0 +1,34 @@
import torch
from torch.distributions.bernoulli import Bernoulli
from .theseus_errors import NoSuccessorError
class TheseusModule(torch.nn.Module):
"""
TheseusModule is the atomic replacing unit.
"""
# TheseusModule will do nothing unless its replacing_rate is specified
def __init__(self, predecessor: torch.nn.Module = None, successor: torch.nn.Module = None, replacing_rate=0):
super().__init__()
self.predecessor = predecessor
self.successor = successor
self.sampler = Bernoulli(torch.FloatTensor([replacing_rate]))
def forward(self, *args, **kwargs):
if self.successor is None:
raise NoSuccessorError(
"The successor is not specified. In this case, do not call `TheseusModule` directly."
)
return self.sample_and_pass()(*args, **kwargs)
def sample_and_pass(self):
# Always replace when `self.training == False`
# Randomly substitute when `self.training == True`
if not self.training or self.sampler.sample() == 1:
return self.successor
return self.predecessor
def set_replacing_rate(self, replacing_rate):
self.sampler = Bernoulli(torch.FloatTensor([replacing_rate]))
+129
View File
@@ -0,0 +1,129 @@
import unittest
from transformers import is_torch_available
from transformers.testing_utils import require_torch
if is_torch_available():
import torch
from transformers import theseus
@require_torch
class TheseusTest(unittest.TestCase):
# These tests have a very small probability to fail even when the code is correct.
# In this case, please re-run the test.
def test_theseus_module(self):
module_a = torch.nn.Linear(1, 1)
module_b = torch.nn.Linear(1, 2)
theseus_module = theseus.TheseusModule(predecessor=module_a, successor=module_b)
for _ in range(10):
self.assertEqual(theseus_module.sample_and_pass(), module_a)
theseus_module.set_replacing_rate(1)
for _ in range(10):
self.assertEqual(theseus_module.sample_and_pass(), module_b)
theseus_module.set_replacing_rate(0.5)
prd_check, scc_check = False, False
for _ in range(100):
if theseus_module.sample_and_pass() == module_a:
prd_check = True
elif theseus_module.sample_and_pass() == module_b:
scc_check = True
self.assertTrue(prd_check)
self.assertTrue(scc_check)
def test_theseus_list_one_to_one(self):
prd_module_list = torch.nn.ModuleList()
scc_module_list = torch.nn.ModuleList()
theseus_module_list = theseus.TheseusList()
for i in range(10):
prd_module = torch.nn.Linear(1, i)
scc_module = torch.nn.Linear(2, i)
theseus_module = theseus.TheseusModule(predecessor=prd_module, successor=scc_module)
prd_module_list.append(prd_module)
scc_module_list.append(scc_module)
theseus_module_list.append(theseus_module)
for _ in range(10):
sampled_list = theseus_module_list.sample_and_pass()
self.assertEqual(len(sampled_list), 10)
for i in range(10):
self.assertEqual(prd_module_list[i], sampled_list[i])
theseus_module_list.set_replacing_rate(1)
for _ in range(10):
sampled_list = theseus_module_list.sample_and_pass()
self.assertEqual(len(sampled_list), 10)
for i in range(10):
self.assertEqual(scc_module_list[i], sampled_list[i])
theseus_module_list.set_replacing_rate(0.5)
prd_num, scc_num = 0, 0
for _ in range(10):
sampled_list = theseus_module_list.sample_and_pass()
self.assertEqual(len(sampled_list), 10)
for i in range(10):
if scc_module_list[i] == sampled_list[i]:
scc_num += 1
elif prd_module_list[i] == sampled_list[i]:
prd_num += 1
self.assertLess(scc_num / (prd_num + scc_num), 1)
self.assertGreater(scc_num / (prd_num + scc_num), 0)
def test_theseus_list_many_to_one(self):
scc_module_list = torch.nn.ModuleList()
theseus_module_list = theseus.TheseusList()
for i in range(10):
prd_module = torch.nn.ModuleList([torch.nn.Linear(1, i), torch.nn.Linear(2, i)])
scc_module = torch.nn.Linear(3, i)
theseus_module = theseus.TheseusModule(predecessor=prd_module, successor=scc_module)
scc_module_list.append(scc_module)
theseus_module_list.append(theseus_module)
for _ in range(10):
sampled_list = theseus_module_list.sample_and_pass()
self.assertEqual(len(sampled_list), 20)
theseus_module_list.set_replacing_rate(1)
for _ in range(10):
sampled_list = theseus_module_list.sample_and_pass()
self.assertEqual(len(sampled_list), 10)
for i in range(10):
self.assertEqual(scc_module_list[i], sampled_list[i])
theseus_module_list.set_replacing_rate(0.5)
len_sum = 0
for _ in range(100):
sampled_list = theseus_module_list.sample_and_pass()
len_sum += len(sampled_list)
len_avg = len_sum / 100
self.assertGreater(len_avg, 10)
self.assertLess(len_avg, 20)
def test_layerdrop_list(self):
layers = torch.nn.ModuleList([torch.nn.Linear(1, i) for i in range(12)])
drop_layers = theseus.LayerDropList.from_module_list(layers, 0)
for _ in range(10):
self.assertEqual(len(drop_layers.sample_and_pass()), 0)
drop_layers.set_replacing_rate(1)
for _ in range(10):
self.assertEqual(len(drop_layers.sample_and_pass()), 12)
drop_layers.set_replacing_rate(0.5)
total_layers = 0
for _ in range(100):
total_layers += len(drop_layers.sample_and_pass())
avg_layers = total_layers / 100
self.assertLess(avg_layers, 12)
self.assertGreater(avg_layers, 0)
def test_mixout_list(self):
layers = torch.nn.ModuleList([torch.nn.Linear(1, i) for i in range(12)])
mixout_layers = theseus.MixoutList.from_module_list(layers, 0)
for theseus_module in mixout_layers:
self.assertNotEqual(theseus_module.predecessor, theseus_module.successor)
self.assertEqual(theseus_module.predecessor.__class__, theseus_module.successor.__class__)
for _ in range(10):
self.assertEqual(len(mixout_layers.sample_and_pass()), 12)
mixout_layers.set_replacing_rate(1)
for _ in range(10):
self.assertEqual(len(mixout_layers.sample_and_pass()), 12)