Switch prd & scc for LayerDrop & mixout

This commit is contained in:
jetrunner
2020-09-17 22:38:10 +08:00
parent 181b4c28a8
commit 20dd34f887
4 changed files with 16 additions and 6 deletions
+1 -1
View File
@@ -11,5 +11,5 @@ class LayerDropList(TheseusList):
def from_module_list(cls, module_list, replacing_rate):
list_to_return = cls()
for module in module_list:
list_to_return.append(TheseusModule(predecessor=module, replacing_rate=replacing_rate))
list_to_return.append(TheseusModule(successor=module, replacing_rate=replacing_rate))
return list_to_return
+12 -2
View File
@@ -10,10 +10,20 @@ class MixoutList(TheseusList):
"""
@classmethod
def from_module_list(cls, module_list, replacing_rate):
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=module, successor=deepcopy(module), replacing_rate=replacing_rate)
TheseusModule(predecessor=predecessor, successor=module, replacing_rate=replacing_rate)
)
return list_to_return
+1 -1
View File
@@ -10,7 +10,7 @@ class TheseusModule(torch.nn.Module):
"""
# TheseusModule will do nothing unless its replacing_rate is specified
def __init__(self, predecessor: torch.nn.Module, successor: torch.nn.Module = None, replacing_rate=0):
def __init__(self, predecessor: torch.nn.Module = None, successor: torch.nn.Module = None, replacing_rate=0):
super().__init__()
self.predecessor = predecessor
self.successor = successor
+2 -2
View File
@@ -104,10 +104,10 @@ class TheseusTest(unittest.TestCase):
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()), 12)
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()), 0)
self.assertEqual(len(drop_layers.sample_and_pass()), 12)
drop_layers.set_replacing_rate(0.5)
total_layers = 0
for _ in range(100):