Team Ai
Datasetpublic

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.

sourceHugging Facemitupdated 10mo agoView on Hugging Face
0likes1.3kdownloads
example.py273 linesDownload Raw Back to root
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