Compare commits
10
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
20dd34f887 | ||
|
|
181b4c28a8 | ||
|
|
3cf1064932 | ||
|
|
d0c53ad1f6 | ||
|
|
33c569f0f1 | ||
|
|
d66af94fa4 | ||
|
|
4cfafd6bdf | ||
|
|
b73d1e6e7f | ||
|
|
fb8dabc2dd | ||
|
|
934139678b |
@@ -0,0 +1,28 @@
|
||||
# transformers.theseus
|
||||

|
||||
|
||||
`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}
|
||||
}
|
||||
```
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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]))
|
||||
@@ -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)
|
||||
Reference in New Issue
Block a user