(decoder, path)
| 44 | |
| 45 | |
| 46 | def save_decoder(decoder, path): |
| 47 | pathlib.Path(path).mkdir(parents=True, exist_ok=True) |
| 48 | |
| 49 | save_conv2d(decoder.conv_in, pathlib.Path(path, 'conv_in')) |
| 50 | save_mid(decoder.mid, pathlib.Path(path, 'mid')) |
| 51 | |
| 52 | for i, block in enumerate(decoder.up[::-1]): |
| 53 | print(i) |
| 54 | if isinstance(block['block'][0].nin_shortcut, Conv2d): |
| 55 | print(block['block'][0].nin_shortcut.weight.shape) |
| 56 | save_decoder_block(block, pathlib.Path(path, f'blocks/{i}')) |
| 57 | |
| 58 | save_scalar(len(decoder.up), "n_block", path) |
| 59 | save_group_norm(decoder.norm_out, pathlib.Path(path, 'norm_out')) |
| 60 | save_conv2d(decoder.conv_out, pathlib.Path(path, 'conv_out')) |
| 61 | |
| 62 | def save_encoder_block(encoder_block, path): |
| 63 | pathlib.Path(path).mkdir(parents=True, exist_ok=True) |
no test coverage detected