Spaces:
Sleeping
Sleeping
File size: 7,608 Bytes
a95f6c0 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 | def load_state_dict( state_dict, network=None, ema=None, optimizer=None, log=True):
'''
utility for loading state dicts for different models. This function sequentially tries different strategies
args:
state_dict: the state dict to load
returns:
True if the state dict was loaded, False otherwise
Assuming the operations are don in_place, this function will not create a copy of the network and optimizer (I hope)
'''
#print(state_dict)
if log: print("Loading state dict")
if log:
print(state_dict.keys())
#if there
try:
if log: print("Attempt 1: trying with strict=True")
if network is not None:
network.load_state_dict(state_dict['network'])
if optimizer is not None:
optimizer.load_state_dict(state_dict['optimizer'])
if ema is not None:
ema.load_state_dict(state_dict['ema'])
return True
except Exception as e:
if log:
print("Could not load state dict")
print(e)
try:
if log: print("Attempt 2: trying with strict=False")
if network is not None:
network.load_state_dict(state_dict['network'], strict=False)
#we cannot load the optimizer in this setting
#self.optimizer.load_state_dict(state_dict['optimizer'], strict=False)
if ema is not None:
ema.load_state_dict(state_dict['ema'], strict=False)
return True
except Exception as e:
if log:
print("Could not load state dict")
print(e)
print("training from scratch")
try:
if log: print("Attempt 3: trying with strict=False,but making sure that the shapes are fine")
if ema is not None:
ema_state_dict = ema.state_dict()
if network is not None:
network_state_dict = network.state_dict()
i=0
if network is not None:
for name, param in state_dict['network'].items():
if log: print("checking",name)
if name in network_state_dict.keys():
if network_state_dict[name].shape==param.shape:
network_state_dict[name]=param
if log:
print("assigning",name)
i+=1
network.load_state_dict(network_state_dict)
if ema is not None:
for name, param in state_dict['ema'].items():
if log: print("checking",name)
if name in ema_state_dict.keys():
if ema_state_dict[name].shape==param.shape:
ema_state_dict[name]=param
if log:
print("assigning",name)
i+=1
ema.load_state_dict(ema_state_dict)
if i==0:
if log: print("WARNING, no parameters were loaded")
raise Exception("No parameters were loaded")
elif i>0:
if log: print("loaded", i, "parameters")
return True
except Exception as e:
print(e)
print("the second strict=False failed")
try:
if log: print("Attempt 4: Assuming the naming is different, with the network and ema called 'state_dict'")
if network is not None:
network.load_state_dict(state_dict['state_dict'])
if ema is not None:
ema.load_state_dict(state_dict['state_dict'])
except Exception as e:
if log:
print("Could not load state dict")
print(e)
print("training from scratch")
print("It failed 3 times!! but not giving up")
#print the names of the parameters in self.network
try:
if log: print("Attempt 5: trying to load with different names, now model='model' and ema='ema_weights'")
if ema is not None:
dic_ema = {}
for (key, tensor) in zip(state_dict['model'].keys(), state_dict['ema_weights']):
dic_ema[key] = tensor
ema.load_state_dict(dic_ema)
return True
except Exception as e:
if log:
print(e)
try:
if log: print("Attempt 6: If there is something wrong with the name of the ema parameters, we can try to load them using the names of the parameters in the model")
if ema is not None:
dic_ema = {}
i=0
for (key, tensor) in zip(state_dict['model'].keys(), state_dict['model'].values()):
if tensor.requires_grad:
dic_ema[key]=state_dict['ema_weights'][i]
i=i+1
else:
dic_ema[key]=tensor
ema.load_state_dict(dic_ema)
return True
except Exception as e:
if log:
print(e)
try:
#assign the parameters in state_dict to self.network using a for loop
print("Attempt 7: Trying to load the parameters one by one. This is for the dance diffusion model, looking for parameters starting with 'diffusion.' or 'diffusion_ema.'")
if ema is not None:
ema_state_dict = ema.state_dict()
if network is not None:
network_state_dict = ema.state_dict()
i=0
if network is not None:
for name, param in state_dict['state_dict'].items():
print("checking",name)
if name.startswith("diffusion."):
i+=1
name=name.replace("diffusion.","")
if network_state_dict[name].shape==param.shape:
#print(param.shape, network.state_dict()[name].shape)
network_state_dict[name]=param
#print("assigning",name)
network.load_state_dict(network_state_dict, strict=False)
if ema is not None:
for name, param in state_dict['state_dict'].items():
if name.startswith("diffusion_ema."):
i+=1
name=name.replace("diffusion_ema.","")
if ema_state_dict[name].shape==param.shape:
if log:
print(param.shape, ema.state_dict()[name].shape)
ema_state_dict[name]=param
ema.load_state_dict(ema_state_dict, strict=False)
if i==0:
print("WARNING, no parameters were loaded")
raise Exception("No parameters were loaded")
elif i>0:
print("loaded", i, "parameters")
return True
except Exception as e:
if log:
print(e)
if network is not None:
network.load_state_dict(state_dict, strict=True)
if ema is not None:
ema.load_state_dict(state_dict, strict=True)
return True
|