| 2 | |
| 3 | |
| 4 | class WeightModule: |
| 5 | def __init__(self): |
| 6 | self._modules = {} |
| 7 | self._parameters = {} |
| 8 | |
| 9 | def is_empty(self): |
| 10 | return len(self._modules) == 0 and len(self._parameters) == 0 |
| 11 | |
| 12 | def add_module(self, name, module): |
| 13 | self._modules[name] = module |
| 14 | setattr(self, name, module) |
| 15 | |
| 16 | def register_parameter(self, name, param): |
| 17 | self._parameters[name] = param |
| 18 | setattr(self, name, param) |
| 19 | |
| 20 | def load(self, weight_dict): |
| 21 | for _, module in self._modules.items(): |
| 22 | if hasattr(module, "load"): |
| 23 | module.load(weight_dict) |
| 24 | |
| 25 | for _, parameter in self._parameters.items(): |
| 26 | if hasattr(parameter, "load"): |
| 27 | parameter.load(weight_dict) |
| 28 | |
| 29 | def register_diff(self, weight_dict): |
| 30 | for _, module in self._modules.items(): |
| 31 | if hasattr(module, "register_diff"): |
| 32 | module.register_diff(weight_dict) |
| 33 | |
| 34 | for _, parameter in self._parameters.items(): |
| 35 | if hasattr(parameter, "register_diff"): |
| 36 | parameter.register_diff(weight_dict) |
| 37 | |
| 38 | def register_lora(self, weight_dict, strength): |
| 39 | for _, module in self._modules.items(): |
| 40 | if hasattr(module, "register_lora"): |
| 41 | module.register_lora(weight_dict, strength) |
| 42 | |
| 43 | for _, parameter in self._parameters.items(): |
| 44 | if hasattr(parameter, "register_lora"): |
| 45 | parameter.register_lora(weight_dict, strength) |
| 46 | |
| 47 | def update_lora(self, weight_dict, strength): |
| 48 | for _, module in self._modules.items(): |
| 49 | if hasattr(module, "update_lora"): |
| 50 | module.update_lora(weight_dict, strength) |
| 51 | |
| 52 | for _, parameter in self._parameters.items(): |
| 53 | if hasattr(parameter, "update_lora"): |
| 54 | parameter.update_lora(weight_dict, strength) |
| 55 | |
| 56 | def remove_lora(self): |
| 57 | for _, module in self._modules.items(): |
| 58 | if hasattr(module, "remove_lora"): |
| 59 | module.remove_lora() |
| 60 | |
| 61 | for _, parameter in self._parameters.items(): |
no outgoing calls
no test coverage detected