MCPcopy Create free account
hub / github.com/OpenImagingLab/FlashVSR / replace_layer

Method replace_layer

diffsynth/models/flux_controlnet.py:175–206  ·  view source on GitHub ↗
(model)

Source from the content-addressed store, hash-verified

173 self.norm_type, self.scale_grad_by_freq, self.sparse)
174
175 def replace_layer(model):
176 for name, module in model.named_children():
177 if isinstance(module,quantized_layer.QRMSNorm):
178 continue
179 if isinstance(module, torch.nn.Linear):
180 with init_weights_on_device():
181 new_layer = quantized_layer.QLinear(module.in_features,module.out_features)
182 new_layer.weight = module.weight
183 if module.bias is not None:
184 new_layer.bias = module.bias
185 setattr(model, name, new_layer)
186 elif isinstance(module, RMSNorm):
187 if hasattr(module,"quantized"):
188 continue
189 module.quantized= True
190 new_layer = quantized_layer.QRMSNorm(module)
191 setattr(model, name, new_layer)
192 elif isinstance(module,torch.nn.Embedding):
193 rows, cols = module.weight.shape
194 new_layer = quantized_layer.QEmbedding(
195 num_embeddings=rows,
196 embedding_dim=cols,
197 _weight=module.weight,
198 # _freeze=module.freeze,
199 padding_idx=module.padding_idx,
200 max_norm=module.max_norm,
201 norm_type=module.norm_type,
202 scale_grad_by_freq=module.scale_grad_by_freq,
203 sparse=module.sparse)
204 setattr(model, name, new_layer)
205 else:
206 replace_layer(module)
207
208 replace_layer(self)
209

Callers

nothing calls this directly

Calls 1

init_weights_on_deviceFunction · 0.85

Tested by

no test coverage detected