kimura-koya/shape-polygons
Shape Polygons Dataset A synthetic dataset containing 70,000 images of various colored polygons (triangles to octagons) rendered on black backgrounds. Dataset Description This dataset consists of programmatically generated polygon images with full metadata about each shape's properties. It's designed for tasks such as: Shape Classification: Classify polygons by number of vertices (3-8) Regression Tasks: Predict shape properties (size, angle, position, color)… See the full description on the dataset page: https://huggingface.co/datasets/kimura-koya/shape-polygons.
01.3k
1"""2Shape Polygons Dataset - Usage Examples3 4This script demonstrates various ways to load and use the Shape Polygons Dataset.5"""6 7import os8import pandas as pd9import matplotlib.pyplot as plt10from PIL import Image11 12# =============================================================================13# Basic Usage: Loading Metadata and Images14# =============================================================================15 16def load_dataset(data_dir=".", split="train"):17 """Load metadata and return as pandas DataFrame."""18 metadata_path = os.path.join(data_dir, split, "metadata.csv")19 return pd.read_csv(metadata_path)20 21 22def load_image(data_dir, split, filename):23 """Load a single image from the dataset."""24 img_path = os.path.join(data_dir, split, "images", filename)25 return Image.open(img_path)26 27 28# =============================================================================29# Example 1: Explore Dataset Statistics30# =============================================================================31 32def explore_statistics(data_dir="."):33 """Print dataset statistics."""34 print("=" * 50)35 print("Shape Polygons Dataset Statistics")36 print("=" * 50)37 38 for split in ["train", "test"]:39 df = load_dataset(data_dir, split)40 print(f"\n{split.upper()} Split:")41 print(f" Total images: {len(df)}")42 print(f"\n Vertices distribution:")43 for v in range(3, 9):44 count = len(df[df["vertices"] == v])45 print(f" {v} vertices: {count} ({count/len(df)*100:.1f}%)")46 47 print(f"\n Size statistics:")48 print(f" Min: {df['size'].min():.4f}")49 print(f" Max: {df['size'].max():.4f}")50 print(f" Mean: {df['size'].mean():.4f}")51 52 53# =============================================================================54# Example 2: Visualize Sample Images55# =============================================================================56 57def visualize_samples(data_dir=".", n_samples=12, split="train"):58 """Visualize random samples from the dataset."""59 df = load_dataset(data_dir, split)60 samples = df.sample(n=min(n_samples, len(df)))61 62 n_cols = 463 n_rows = (len(samples) + n_cols - 1) // n_cols64 65 fig, axes = plt.subplots(n_rows, n_cols, figsize=(12, 3 * n_rows))66 axes = axes.flatten() if n_samples > 1 else [axes]67 68 for idx, (_, row) in enumerate(samples.iterrows()):69 img = load_image(data_dir, split, row["filename"])70 axes[idx].imshow(img)71 axes[idx].set_title(f"{row['vertices']} vertices\nsize={row['size']:.2f}")72 axes[idx].axis("off")73 74 # Hide empty subplots75 for idx in range(len(samples), len(axes)):76 axes[idx].axis("off")77 78 plt.tight_layout()79 plt.savefig("samples_visualization.png", dpi=150, bbox_inches="tight")80 print(f"Saved visualization to 'samples_visualization.png'")81 plt.show()82 83 84# =============================================================================85# Example 3: Visualize by Shape Type86# =============================================================================87 88def visualize_by_shape_type(data_dir=".", split="train"):89 """Show one example of each shape type."""90 df = load_dataset(data_dir, split)91 shape_names = {92 3: "Triangle",93 4: "Quadrilateral",94 5: "Pentagon",95 6: "Hexagon",96 7: "Heptagon",97 8: "Octagon"98 }99 100 fig, axes = plt.subplots(2, 3, figsize=(12, 8))101 axes = axes.flatten()102 103 for idx, vertices in enumerate(range(3, 9)):104 sample = df[df["vertices"] == vertices].iloc[0]105 img = load_image(data_dir, split, sample["filename"])106 axes[idx].imshow(img)107 axes[idx].set_title(f"{shape_names[vertices]}\n({vertices} vertices)")108 axes[idx].axis("off")109 110 plt.suptitle("Shape Types in Dataset", fontsize=14, fontweight="bold")111 plt.tight_layout()112 plt.savefig("shape_types.png", dpi=150, bbox_inches="tight")113 print(f"Saved visualization to 'shape_types.png'")114 plt.show()115 116 117# =============================================================================118# Example 4: PyTorch Dataset Class119# =============================================================================120 121# Optional imports for PyTorch functionality122try:123 import torch124 from torch.utils.data import Dataset, DataLoader125 from torchvision import transforms126 PYTORCH_AVAILABLE = True127except ImportError:128 PYTORCH_AVAILABLE = False129 130 131class ShapePolygonsDataset:132 """PyTorch Dataset for Shape Polygons.133 134 Requires: torch, torchvision135 """136 137 def __init__(self, root_dir, split="train", transform=None, task="classification"):138 """139 Args:140 root_dir: Root directory of the dataset141 split: "train" or "test"142 transform: Optional torchvision transforms143 task: "classification" for vertex count, "regression" for size prediction, "multi" for all properties144 """145 if not PYTORCH_AVAILABLE:146 raise ImportError("PyTorch is required. Install with: pip install torch torchvision")147 148 self.root_dir = root_dir149 self.split = split150 self.transform = transform151 self.task = task152 self.metadata = pd.read_csv(os.path.join(root_dir, split, "metadata.csv"))153 154 if self.transform is None:155 self.transform = transforms.Compose([156 transforms.ToTensor(),157 transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5])158 ])159 160 def __len__(self):161 return len(self.metadata)162 163 def __getitem__(self, idx):164 row = self.metadata.iloc[idx]165 img_path = os.path.join(self.root_dir, self.split, "images", row["filename"])166 image = Image.open(img_path).convert("RGB")167 168 if self.transform:169 image = self.transform(image)170 171 if self.task == "classification":172 # Class label: 0-5 for 3-8 vertices173 label = torch.tensor(row["vertices"] - 3, dtype=torch.long)174 elif self.task == "regression":175 # Predict size176 label = torch.tensor(row["size"], dtype=torch.float32)177 elif self.task == "multi":178 # Multi-task: return all properties179 label = {180 "vertices": torch.tensor(row["vertices"] - 3, dtype=torch.long),181 "size": torch.tensor(row["size"], dtype=torch.float32),182 "angle": torch.tensor(row["angle"], dtype=torch.float32),183 "center": torch.tensor([row["center_x"], row["center_y"]], dtype=torch.float32),184 "color": torch.tensor([row["color_r"], row["color_g"], row["color_b"]], dtype=torch.float32)185 }186 else:187 raise ValueError(f"Unknown task: {self.task}")188 189 return image, label190 191 192def demo_pytorch_dataloader(data_dir="."):193 """Demonstrate PyTorch DataLoader usage."""194 if not PYTORCH_AVAILABLE:195 print("PyTorch is not installed. Install with: pip install torch torchvision")196 return197 198 print("Creating PyTorch Dataset and DataLoader...")199 dataset = ShapePolygonsDataset(data_dir, split="train", task="classification")200 dataloader = DataLoader(dataset, batch_size=32, shuffle=True, num_workers=0)201 202 # Get one batch203 images, labels = next(iter(dataloader))204 print(f"Batch shape: {images.shape}")205 print(f"Labels shape: {labels.shape}")206 print(f"Label values (vertices - 3): {labels[:10].tolist()}")207 print(f"Actual vertex counts: {[l + 3 for l in labels[:10].tolist()]}")208 209 210# =============================================================================211# Example 5: Color Analysis212# =============================================================================213 214def analyze_colors(data_dir=".", split="train"):215 """Analyze color distribution in the dataset."""216 df = load_dataset(data_dir, split)217 218 fig, axes = plt.subplots(1, 3, figsize=(15, 4))219 colors = ["red", "green", "blue"]220 columns = ["color_r", "color_g", "color_b"]221 222 for idx, (color, col) in enumerate(zip(colors, columns)):223 axes[idx].hist(df[col], bins=50, color=color, alpha=0.7, edgecolor="black")224 axes[idx].set_xlabel(f"{color.capitalize()} Value")225 axes[idx].set_ylabel("Frequency")226 axes[idx].set_title(f"{color.capitalize()} Channel Distribution")227 228 plt.suptitle(f"Color Distribution in {split.capitalize()} Set", fontsize=14, fontweight="bold")229 plt.tight_layout()230 plt.savefig("color_distribution.png", dpi=150, bbox_inches="tight")231 print(f"Saved visualization to 'color_distribution.png'")232 plt.show()233 234 235# =============================================================================236# Main237# =============================================================================238 239if __name__ == "__main__":240 import argparse241 242 parser = argparse.ArgumentParser(description="Shape Polygons Dataset Examples")243 parser.add_argument("--data-dir", type=str, default=".", help="Path to dataset root")244 parser.add_argument(245 "--example",246 type=str,247 choices=["stats", "samples", "shapes", "pytorch", "colors", "all"],248 default="all",249 help="Which example to run"250 )251 args = parser.parse_args()252 253 examples = {254 "stats": ("Dataset Statistics", lambda: explore_statistics(args.data_dir)),255 "samples": ("Sample Visualization", lambda: visualize_samples(args.data_dir)),256 "shapes": ("Shape Types", lambda: visualize_by_shape_type(args.data_dir)),257 "pytorch": ("PyTorch DataLoader Demo", lambda: demo_pytorch_dataloader(args.data_dir)),258 "colors": ("Color Analysis", lambda: analyze_colors(args.data_dir)),259 }260 261 if args.example == "all":262 for name, (desc, func) in examples.items():263 print(f"\n{'=' * 50}")264 print(f"Example: {desc}")265 print("=" * 50)266 func()267 else:268 name = args.example269 desc, func = examples[name]270 print(f"Running Example: {desc}")271 func()272 273 