(
mid_block_type: str,
num_layers: int,
in_channels: int,
mid_channels: int,
out_channels: int,
embed_dim: int,
add_downsample: bool,
)
| 668 | |
| 669 | |
| 670 | def get_mid_block( |
| 671 | mid_block_type: str, |
| 672 | num_layers: int, |
| 673 | in_channels: int, |
| 674 | mid_channels: int, |
| 675 | out_channels: int, |
| 676 | embed_dim: int, |
| 677 | add_downsample: bool, |
| 678 | ) -> MidBlockType: |
| 679 | if mid_block_type == "MidResTemporalBlock1D": |
| 680 | return MidResTemporalBlock1D( |
| 681 | num_layers=num_layers, |
| 682 | in_channels=in_channels, |
| 683 | out_channels=out_channels, |
| 684 | embed_dim=embed_dim, |
| 685 | add_downsample=add_downsample, |
| 686 | ) |
| 687 | elif mid_block_type == "ValueFunctionMidBlock1D": |
| 688 | return ValueFunctionMidBlock1D(in_channels=in_channels, out_channels=out_channels, embed_dim=embed_dim) |
| 689 | elif mid_block_type == "UNetMidBlock1D": |
| 690 | return UNetMidBlock1D(in_channels=in_channels, mid_channels=mid_channels, out_channels=out_channels) |
| 691 | raise ValueError(f"{mid_block_type} does not exist.") |
| 692 | |
| 693 | |
| 694 | def get_out_block( |
no test coverage detected
searching dependent graphs…