Return the targeted layer (nn.Module Object) given a hierarchical layer name, separated by /. Args: model (model): model to get layers from. layer_name (str): name of the layer. Returns: prev_module (nn.Module): the layer from the model with `layer_name` name
(model, layer_name)
| 323 | |
| 324 | |
| 325 | def get_layer(model, layer_name): |
| 326 | """ |
| 327 | Return the targeted layer (nn.Module Object) given a hierarchical layer name, |
| 328 | separated by /. |
| 329 | Args: |
| 330 | model (model): model to get layers from. |
| 331 | layer_name (str): name of the layer. |
| 332 | Returns: |
| 333 | prev_module (nn.Module): the layer from the model with `layer_name` name. |
| 334 | """ |
| 335 | layer_ls = layer_name.split("/") |
| 336 | prev_module = model |
| 337 | for layer in layer_ls: |
| 338 | prev_module = prev_module._modules[layer] |
| 339 | |
| 340 | return prev_module |
| 341 | |
| 342 | |
| 343 | class TaskInfo: |
no outgoing calls
no test coverage detected