(
name="CircularConv2D",
**kwargs
)
| 107 | |
| 108 | |
| 109 | def get_module( |
| 110 | name="CircularConv2D", |
| 111 | **kwargs |
| 112 | ): |
| 113 | if name == "CircularConv2D": |
| 114 | return CircularConv2D(**kwargs) |
| 115 | elif name == "Conv2D": |
| 116 | return torch.nn.Conv2d(**kwargs) |
| 117 | elif name == "Conv2DSiLU": |
| 118 | return Conv2DSiLU(**kwargs) |
| 119 | elif name == "ResBlock": |
| 120 | return ResnetBlock(**kwargs) |
| 121 | elif name == "ResConv2DBlock": |
| 122 | return ResnetConv2DBlock(**kwargs) |
| 123 | elif name == "Upsample": |
| 124 | return Upsample(**kwargs) |
| 125 | elif name == "Downsample": |
| 126 | return Downsample(**kwargs) |
| 127 | elif name == "Attention": |
| 128 | return make_attn(**kwargs) |
| 129 | # ---- module ---- |
| 130 | |
| 131 | # ---- Directional Position Encoding ---- |
no test coverage detected