(model, batch, out_file)
| 55 | |
| 56 | |
| 57 | def reconstruct(model, batch, out_file): |
| 58 | # Reconstruct a single batch only |
| 59 | images = mx.array(batch["image"]) |
| 60 | images_recon = model(images)[0] |
| 61 | paired_images = mx.stack([images, images_recon]).swapaxes(0, 1).flatten(0, 1) |
| 62 | grid_image = grid_image_from_batch(paired_images, num_rows=16) |
| 63 | grid_image.save(out_file) |
| 64 | |
| 65 | |
| 66 | def generate( |
no test coverage detected