Updates paths inside resnets to the new naming scheme (local renaming)
(old_list, n_shave_prefix_segments=0, is_temporal=False)
| 554 | |
| 555 | |
| 556 | def renew_vae_resnet_paths(old_list, n_shave_prefix_segments=0, is_temporal=False): |
| 557 | """ |
| 558 | Updates paths inside resnets to the new naming scheme (local renaming) |
| 559 | """ |
| 560 | mapping = [] |
| 561 | for old_item in old_list: |
| 562 | new_item = old_item |
| 563 | |
| 564 | # Temporal resnet |
| 565 | new_item = old_item.replace("in_layers.0", "norm1") |
| 566 | new_item = new_item.replace("in_layers.2", "conv1") |
| 567 | |
| 568 | new_item = new_item.replace("out_layers.0", "norm2") |
| 569 | new_item = new_item.replace("out_layers.3", "conv2") |
| 570 | |
| 571 | new_item = new_item.replace("skip_connection", "conv_shortcut") |
| 572 | |
| 573 | new_item = new_item.replace("time_stack.", "temporal_res_block.") |
| 574 | |
| 575 | # Spatial resnet |
| 576 | new_item = new_item.replace("conv1", "spatial_res_block.conv1") |
| 577 | new_item = new_item.replace("norm1", "spatial_res_block.norm1") |
| 578 | |
| 579 | new_item = new_item.replace("conv2", "spatial_res_block.conv2") |
| 580 | new_item = new_item.replace("norm2", "spatial_res_block.norm2") |
| 581 | |
| 582 | new_item = new_item.replace("nin_shortcut", "spatial_res_block.conv_shortcut") |
| 583 | |
| 584 | new_item = new_item.replace("mix_factor", "spatial_res_block.time_mixer.mix_factor") |
| 585 | |
| 586 | new_item = shave_segments(new_item, n_shave_prefix_segments=n_shave_prefix_segments) |
| 587 | |
| 588 | mapping.append({"old": old_item, "new": new_item}) |
| 589 | |
| 590 | return mapping |
| 591 | |
| 592 | |
| 593 | def renew_vae_attention_paths(old_list, n_shave_prefix_segments=0): |
no test coverage detected