(self, unlocked_layers:int=0, freeze_layer_norm:bool=True)
| 218 | return self.proj(pooled_out) |
| 219 | |
| 220 | def lock(self, unlocked_layers:int=0, freeze_layer_norm:bool=True): |
| 221 | if not unlocked_layers: # full freezing |
| 222 | for n, p in self.transformer.named_parameters(): |
| 223 | p.requires_grad = (not freeze_layer_norm) if "LayerNorm" in n.split(".") else False |
| 224 | return |
| 225 | |
| 226 | encoder = self.transformer.encoder if hasattr(self.transformer, 'encoder') else self.transformer |
| 227 | layer_list = getattr(encoder, arch_dict[self.config.model_type]["config_names"]["layer_attr"]) |
| 228 | print(f"Unlocking {unlocked_layers}/{len(layer_list) + 1} layers of hf model") |
| 229 | embeddings = getattr( |
| 230 | self.transformer, arch_dict[self.config.model_type]["config_names"]["token_embeddings_attr"]) |
| 231 | modules = [embeddings, *layer_list][:-unlocked_layers] |
| 232 | # freeze layers |
| 233 | for module in modules: |
| 234 | for n, p in module.named_parameters(): |
| 235 | p.requires_grad = (not freeze_layer_norm) if "LayerNorm" in n.split(".") else False |
| 236 | |
| 237 | |
| 238 | @torch.jit.ignore |
nothing calls this directly
no outgoing calls
no test coverage detected