()
| 76 | ) |
| 77 | |
| 78 | def main(): |
| 79 | parser = argparse.ArgumentParser(description="Visualize video sequence batches") |
| 80 | parser.add_argument("--dataset", type=str, default="SONIC", help="Dataset to use") |
| 81 | parser.add_argument("--batch_size", type=int, default=4, help="Batch size") |
| 82 | parser.add_argument("--context_length", type=int, default=4, help="Context length") |
| 83 | parser.add_argument("--num_batches", type=int, default=4, help="Number of batches to visualize") |
| 84 | parser.add_argument("--save_dir", type=str, default="batch_visualizations", help="Directory to save visualizations") |
| 85 | parser.add_argument("--max_batch_display", type=int, default=8, help="Maximum batch elements to display") |
| 86 | parser.add_argument("--max_seq_display", type=int, default=8, help="Maximum timesteps to display") |
| 87 | args = parser.parse_args() |
| 88 | |
| 89 | print(f"Loading {args.dataset} dataset...") |
| 90 | |
| 91 | # load data |
| 92 | _, _, validation_loader, _, _ = load_data_and_data_loaders( |
| 93 | dataset=args.dataset, |
| 94 | batch_size=args.batch_size, |
| 95 | num_frames=args.context_length |
| 96 | ) |
| 97 | # visualize batches |
| 98 | for batch_idx, (frames, _) in enumerate(validation_loader): |
| 99 | if batch_idx >= args.num_batches: |
| 100 | break |
| 101 | |
| 102 | # calculate statistics |
| 103 | frames_cpu = frames.detach().cpu() |
| 104 | batch_size, seq_len, C, H, W = frames_cpu.shape |
| 105 | |
| 106 | # visualize batch |
| 107 | os.makedirs(args.save_dir, exist_ok=True) |
| 108 | save_path = os.path.join(args.save_dir, f"{args.dataset}_batch_{batch_idx + 1}.png") |
| 109 | visualize_batch( |
| 110 | frames, |
| 111 | save_path=save_path, |
| 112 | title=f"Batch {batch_idx + 1} - {args.dataset} Sequences", |
| 113 | max_batch_size=args.max_batch_display, |
| 114 | max_seq_length=args.max_seq_display |
| 115 | ) |
| 116 | |
| 117 | if __name__ == "__main__": |
| 118 | main() |
no test coverage detected