MCPcopy Create free account
hub / github.com/AlmondGod/tinyworlds / main

Function main

scripts/visualize_batch.py:78–115  ·  view source on GitHub ↗
()

Source from the content-addressed store, hash-verified

76 )
77
78def 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
117if __name__ == "__main__":
118 main()

Callers 1

visualize_batch.pyFile · 0.70

Calls 2

visualize_batchFunction · 0.85

Tested by

no test coverage detected