(path: &str, device: &B::Device)
| 133 | } |
| 134 | |
| 135 | fn load_decoder<B: Backend>(path: &str, device: &B::Device) -> Result<Decoder<B>, Box<dyn Error>> { |
| 136 | let conv_in = load_conv2d(&format!("{}/{}", path, "conv_in"), device)?; |
| 137 | let mid = load_mid(&format!("{}/{}", path, "mid"), device)?; |
| 138 | |
| 139 | let n_block = load_usize::<B>("n_block", path, device)?; |
| 140 | let mut blocks = (0..n_block) |
| 141 | .into_iter() |
| 142 | .map(|i| load_decoder_block::<B>(&format!("{}/blocks/{}", path, i), device)) |
| 143 | .collect::<Result<Vec<_>, _>>()?; |
| 144 | |
| 145 | let norm_out = load_group_norm(&format!("{}/{}", path, "norm_out"), device)?; |
| 146 | let silu = SILU {}; |
| 147 | let conv_out = load_conv2d(&format!("{}/{}", path, "conv_out"), device)?; |
| 148 | |
| 149 | Ok(Decoder { |
| 150 | conv_in, |
| 151 | mid, |
| 152 | blocks, |
| 153 | norm_out, |
| 154 | silu, |
| 155 | conv_out, |
| 156 | }) |
| 157 | } |
| 158 | |
| 159 | fn load_encoder<B: Backend>(path: &str, device: &B::Device) -> Result<Encoder<B>, Box<dyn Error>> { |
| 160 | let conv_in = load_conv2d(&format!("{}/{}", path, "conv_in"), device)?; |
no test coverage detected