(self, unlocked_layers: int = 0, freeze_layer_norm: bool = True)
| 169 | return projected |
| 170 | |
| 171 | def lock(self, unlocked_layers: int = 0, freeze_layer_norm: bool = True): |
| 172 | if not unlocked_layers: # full freezing |
| 173 | for n, p in self.transformer.named_parameters(): |
| 174 | p.requires_grad = (not freeze_layer_norm) if "LayerNorm" in n.split(".") else False |
| 175 | return |
| 176 | |
| 177 | encoder = self.transformer.encoder if hasattr(self.transformer, 'encoder') else self.transformer |
| 178 | layer_list = getattr(encoder, arch_dict[self.config.model_type]["config_names"]["layer_attr"]) |
| 179 | print(f"Unlocking {unlocked_layers}/{len(layer_list) + 1} layers of hf model") |
| 180 | embeddings = getattr( |
| 181 | self.transformer, arch_dict[self.config.model_type]["config_names"]["token_embeddings_attr"]) |
| 182 | modules = [embeddings, *layer_list][:-unlocked_layers] |
| 183 | # freeze layers |
| 184 | for module in modules: |
| 185 | for n, p in module.named_parameters(): |
| 186 | p.requires_grad = (not freeze_layer_norm) if "LayerNorm" in n.split(".") else False |
| 187 | |
| 188 | @torch.jit.ignore |
| 189 | def set_grad_checkpointing(self, enable=True): |
nothing calls this directly
no outgoing calls
no test coverage detected