Switch prd & scc for LayerDrop & mixout
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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):
|
||||
|
||||
Reference in New Issue
Block a user