Skip to content

Instantly share code, notes, and snippets.

@ottobricks
Last active June 15, 2026 16:01
Show Gist options
  • Select an option

  • Save ottobricks/158ce18e90eef164e1dc80cfd1caecea to your computer and use it in GitHub Desktop.

Select an option

Save ottobricks/158ce18e90eef164e1dc80cfd1caecea to your computer and use it in GitHub Desktop.
pytorch_binary_classification.py
# /// script
# requires-python = ">=3.11"
# dependencies = [
# "pandas",
# "torch",
# "torchvision",
# "scikit-learn"
# ]
# ///
import random
from typing import Literal
from pandas import DataFrame
from sklearn.model_selection import train_test_split
from torch.utils.data import DataLoader, Dataset, Subset, random_split
from torchvision import datasets, transforms
class CIFAKEImageDataset(Dataset):
def __init__(self, dataframe):
self.dataframe = dataframe
def transform():
raise NotImplementedError("You must implement method 'transform'")
def __len__(self):
return len(self.dataframe)
def __getitem__(self, idx):
image = self.dataframe.iloc[idx]["image"]
label = self.dataframe.iloc[idx]["label"]
image = image.convert("RGB")
if self.transform:
image = self.transform(image)
return image, label
def main() -> None:
data = get_data()
train_dataframe, test_dataframe = train_test_split(
data, test_size=0.2, random_state=42, stratify=data["label"]
)
train_dataset = CIFAKEImageDataset(train_dataframe)
test_dataset = CIFAKEImageDataset(test_dataframe)
train_loader = DataLoader(train_dataset, batch_size=10, shuffle=False)
test_loader = DataLoader(test_dataset, batch_size=10, shuffle=False)
def get_data(split: Literal["train", "test"] = "train", sample_size:int = 100) -> DataFrame:
dataset = load_dataset("dragonintelligence/CIFAKE-image-dataset", split=split)
sample_dataset = dataset.select(random.sample(range(len(dataset)), sample_size))
return DataFrame(
[
{
"label": item["label"],
"label_text": "real" if item["label"] == 0 else "fake",
"image": item["image"],
}
for item in sample_dataset
]
)
if __name__ == "__main__":
main()
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment