{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[],"dockerImageVersionId":28755,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# [Perceiver: General Perception with Iterative Attention](https://arxiv.org/pdf/2103.03206)","metadata":{}},{"cell_type":"code","source":"!pip uninstall torch -y\n!pip uninstall torchvision torchaudio -y\n!pip install torch==2.9.0 torch-xla==2.9.0\n!pip install torchvision==0.24.0\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nimport torch.optim as optim\nimport numpy as np","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from dataclasses import dataclass\n\n@dataclass\nclass Config:\n    data_dir: str = \"/kaggle/input/competitions/imagenet-object-localization-challenge/ILSVRC/Data/CLS-LOC\"\n    image_size: int = 224\n    num_classes: int = 1000\n\n    # Model\n    latent_dim: int = 512\n    num_latents: int = 512\n    num_blocks: int = 8\n    num_heads: int = 8\n    head_dim: int = 64\n    num_bands: int = 64\n    max_resolution: int = 224\n    chunk_size: int = 2048\n\n    # Training\n    batch_size_per_core: int = 4\n    effective_batch_size: int = 512\n    num_epochs: int = 2\n    warmup_epochs: int = 10\n    base_lr: float = 5e-4\n    weight_decay: float = 0.01\n    adam_beta1: float = 0.9\n    adam_beta2: float = 0.999\n    label_smoothing: float = 0.1\n\n    # System\n    num_workers: int = 4\n    pin_memory: bool = True\n    subset_fraction: float = 0.01 \n    cache_dir: str = \"./dataset_cache\"\n    \n# Create config\nconfig = Config()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from pathlib import Path\nimport pandas as pd\nimport os\nfrom torchvision.datasets import ImageFolder\nfrom torchvision import transforms\nfrom PIL import Image\n\nclass ImageNetValDataset(Dataset):\n    def __init__(self, img_dir, label_csv, class_to_idx, transform=None):\n        self.img_dir = Path(img_dir)\n        self.transform = transform\n        self.class_to_idx = class_to_idx\n\n        df = pd.read_csv(label_csv)\n        self.images = []\n        self.labels = []\n        for _, row in df.iterrows():\n            img_id = row['ImageId']\n            pred_str = row['PredictionString']\n            class_name = pred_str.split()[0]\n            if class_name in self.class_to_idx:\n                self.images.append(img_id + '.JPEG')\n                self.labels.append(class_to_idx[class_name])\n        self.images = np.array(self.images)\n        self.labels = np.array(self.labels)\n\n        H, W = 224, 224\n        self.coords = self._build_coord_grids(H, W)\n\n    @classmethod\n    def from_cached(cls, img_dir, images, labels, class_to_idx, transform):\n        obj = cls.__new__(cls)\n        obj.img_dir = Path(img_dir)\n        obj.transform = transform\n        obj.class_to_idx = class_to_idx\n        obj.images = images\n        obj.labels = labels\n        H = W = 224\n        xs = torch.linspace(-1, 1, steps=W)\n        ys = torch.linspace(-1, 1, steps=H)\n        yv, xv = torch.meshgrid(ys, xs, indexing='ij')\n        obj.coords = torch.stack([xv, yv], dim=-1).view(-1, 2)\n        return obj\n\n    def _build_coord_grids(self, H, W):\n        xs = torch.linspace(-1, 1, W)\n        ys = torch.linspace(-1, 1, H)\n        yv, xv = torch.meshgrid(ys, xs, indexing='ij')\n        return torch.stack([xv, yv], dim=-1).view(-1, 2)\n        \n    def __len__(self):\n        return len(self.images)\n\n    def __getitem__(self, idx):\n        img_path = self.img_dir / self.images[idx]\n        img = Image.open(img_path).convert('RGB')\n        if self.transform:\n            img = self.transform(img)\n        else:\n            img = transforms.PILToTensor()(img)\n\n        img = img.float().permute(1,2,0).reshape(-1, 3)\n        \n        return img, self.coords, torch.tensor(self.labels[idx])","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class CachedSubset(Dataset):\n    def __init__(self, subset, class_to_idx, transform=None):\n        self.subset = subset\n        self.class_to_idx = class_to_idx\n        self.transform = transform\n        H = W = 224\n        xs = torch.linspace(-1, 1, steps=W)\n        ys = torch.linspace(-1, 1, steps=H)\n        yv, xv = torch.meshgrid(ys, xs, indexing='ij')\n        self.coords = torch.stack([xv, yv], dim=-1).view(-1, 2)\n\n    def __len__(self):\n        return len(self.subset)\n\n    def __getitem__(self, idx):\n        img, label = self.subset[idx]\n        if self.transform:\n            img = self.transform(img)\n        else:\n            if not isinstance(img, torch.Tensor):\n                img = transforms.PILToTensor()(img)\n        img = img.float().permute(1, 2, 0).reshape(-1, 3)\n        return img, self.coords, label","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pickle\nimport json\n\ndef build_dataloaders(config):\n    cache_dir = Path(getattr(config, \"cache_dir\", \"./dataset_cache\"))\n    cache_dir.mkdir(exist_ok=True)\n\n    train_indices_file = cache_dir / 'train_indices.npy'\n    class_to_idx_file = cache_dir / 'class_to_idx.json'\n    val_images_file = cache_dir / 'val_images.npy'\n    val_labels_file = cache_dir / 'val_labels.npy'\n\n    train_transform = transforms.Compose([\n        transforms.RandomResizedCrop(224),\n        transforms.RandomHorizontalFlip(),\n        transforms.PILToTensor()\n    ])\n    val_transform = transforms.Compose([\n        transforms.Resize(256),\n        transforms.CenterCrop(224),\n        transforms.PILToTensor()\n    ])\n\n    if train_indices_file.exists() and class_to_idx_file.exists():\n        train_indices = np.load(train_indices_file)\n        with open(class_to_idx_file, 'r') as f:\n            class_to_idx = json.load(f)\n        full_train = ImageFolder(root=os.path.join(config.data_dir, 'train'), transform=None)\n        train_dataset = torch.utils.data.Subset(full_train, train_indices)\n        train_dataset = CachedSubset(train_dataset, class_to_idx, transform=train_transform)\n    else:\n        full_train = ImageFolder(root=os.path.join(config.data_dir, 'train'), transform=None)\n        class_to_idx = full_train.class_to_idx\n        if config.subset_fraction < 1.0:\n            num_samples = int(len(full_train) * config.subset_fraction)\n            indices = np.random.choice(len(full_train), num_samples, replace=False)\n            np.save(train_indices_file, indices)\n            train_dataset = torch.utils.data.Subset(full_train, indices)\n        else:\n            indices = np.arange(len(full_train))\n            np.save(train_indices_file, indices)\n            train_dataset = full_train\n        with open(class_to_idx_file, 'w') as f:\n            json.dump(class_to_idx, f)\n        train_dataset = CachedSubset(train_dataset, class_to_idx, transform=train_transform)\n\n    if val_images_file.exists() and val_labels_file.exists():\n        images = np.load(val_images_file)\n        labels = np.load(val_labels_file)\n        val_dataset = ImageNetValDataset.from_cached(\n            img_dir=os.path.join(config.data_dir, 'val'),\n            images=images,\n            labels=labels,\n            class_to_idx=class_to_idx,\n            transform=val_transform\n        )\n    else:\n        val_dataset = ImageNetValDataset(\n            img_dir=os.path.join(config.data_dir, 'val'),\n            label_csv='/kaggle/input/competitions/imagenet-object-localization-challenge/LOC_val_solution.csv',\n            class_to_idx=class_to_idx,\n            transform=val_transform\n        )\n        np.save(val_images_file, val_dataset.images)\n        np.save(val_labels_file, val_dataset.labels)\n\n    train_loader = DataLoader(train_dataset, batch_size=config.batch_size_per_core,\n                              shuffle=True, num_workers=config.num_workers, drop_last=True)\n    val_loader = DataLoader(val_dataset, batch_size=config.batch_size_per_core * 2,\n                            shuffle=False, num_workers=config.num_workers)\n    \n    return train_loader, val_loader, len(class_to_idx)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class FourierEncode(nn.Module):\n    def __init__(self, num_bands=64, max_resolution=224):\n        super().__init__()\n        self.frequency_band = torch.linspace(1, max_resolution/2, num_bands) # nyquist criterion\n        self.register_buffer('freqs', self.frequency_band)\n\n    def forward(self, x):\n        x_proj = torch.pi * x.unsqueeze(-1) * self.freqs\n        sin_feat = torch.sin(x_proj)\n        cos_feat = torch.cos(x_proj)\n        encoded = torch.stack([sin_feat, cos_feat], dim=-1).flatten(-3)\n        return encoded","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def attention(Q, K, V):\n    d_k = Q.size(-1)\n    scores = torch.matmul(Q, K.transpose(-2, -1)) / np.sqrt(d_k)\n    attn_weights = F.softmax(scores, dim=-1)\n    return torch.matmul(attn_weights, V)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# FOR GPU\n\n# class CrossAttention(nn.Module):\n#     def __init__(self, dim_q, dim_kv, num_heads=8, chunk_size=2048):\n#         super().__init__()\n#         self.num_heads = num_heads\n#         self.chunk_size = chunk_size\n#         self.d_k = dim_q // num_heads\n#         self.q = nn.Linear(dim_q, dim_q, bias=False)\n#         self.k = nn.Linear(dim_kv, dim_q, bias=False)\n#         self.v = nn.Linear(dim_kv, dim_q, bias=False)\n#         self.W_o = nn.Linear(dim_q, dim_q)\n\n#     def forward(self, x_q, x_kv):\n#         B, N, D = x_q.shape # x_q (B, N, D) latents\n#         M = x_kv.shape[1] # x_kv (B, M, C_in) byte inputs\n        \n#         Q = self.q(x_q).reshape(B, N, self.num_heads, self.d_k).permute(0, 2, 1, 3)\n#         K = self.k(x_kv).reshape(B, N, self.num_heads, self.d_k).permute(0, 2, 1, 3)\n#         V = self.v(x_kv).reshape(B, N, self.num_heads, self.d_k).permute(0, 2, 1, 3)\n        \n#         global_sum = torch.zeros((B, self.num_heads, N, 1), device=Q.device)\n#         global_max = torch.full((B, self.num_heads, N, 1), -float('inf'), device=Q.device)\n#         out = torch.zeros_like(Q)\n\n#         for i in range(0, M, self.chunk_size):\n#             K_chunk = K[: ,: ,i:i+self.chunk_size, :]\n#             V_chunk = V[: ,: ,i:i+self.chunk_size, :]\n#             scores = torch.matmul(Q, K_chunk.transpose(-1, -2)) / (self.d_k ** 0.5)\n#             chunk_max = scores.max(dim=-1, keepdim=True)[0]\n#             global_max = torch.maximum(chunk_max, global_max)\n\n#         for i in range(0, M, self.chunk_size):\n#             K_chunk = K[: ,: ,i:i+self.chunk_size, :]\n#             V_chunk = V[: ,: ,i:i+self.chunk_size, :]\n#             scores = torch.matmul(Q, K_chunk.transpose(-1, -2)) / (self.d_k ** 0.5)\n#             exp_scores = torch.exp(scores - global_max)\n#             global_sum += exp_scores.sum(dim=-1, keepdim=True)\n#             out += torch.matmul(exp_scores, V_chunk)\n            \n#         out = out / global_sum\n#         out = out.permute(0, 2, 1 ,3).reshape(B, N, D)\n#         return self.W_o(out)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# FOR TPU\n\nclass CrossAttention(nn.Module):\n    def __init__(self, dim_q, dim_kv, num_heads=8, chunk_size=2048):\n        super().__init__()\n        self.num_heads = num_heads\n        self.head_dim = dim_q // num_heads\n        self.scale = self.head_dim ** -0.5\n        self.to_q = nn.Linear(dim_q, dim_q, bias=False)\n        self.to_k = nn.Linear(dim_kv, dim_q, bias=False)\n        self.to_v = nn.Linear(dim_kv, dim_q, bias=False)\n        self.to_out = nn.Linear(dim_q, dim_q)\n\n    def forward(self, x_q, x_kv):\n        B, N, D = x_q.shape\n        M = x_kv.shape[1]\n\n        Q = self.to_q(x_q).view(B, N, self.num_heads, self.head_dim).permute(0,2,1,3)\n        K = self.to_k(x_kv).view(B, M, self.num_heads, self.head_dim).permute(0,2,1,3)\n        V = self.to_v(x_kv).view(B, M, self.num_heads, self.head_dim).permute(0,2,1,3)\n\n        attn_scores = torch.matmul(Q, K.transpose(-2, -1)) * self.scale\n        attn_weights = F.softmax(attn_scores, dim=-1)\n        out = torch.matmul(attn_weights, V)\n        out = out.permute(0,2,1,3).contiguous().view(B, N, D)\n        return self.to_out(out)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class SelfAttention(nn.Module):\n    def __init__(self, dim, num_heads=8):\n        super().__init__()\n        self.num_heads = num_heads\n        self.head_dim = dim // num_heads\n        self.scale = self.head_dim ** -0.5\n        self.to_qkv = nn.Linear(dim, dim * 3, bias=False)\n        self.to_out = nn.Linear(dim, dim)\n        \n    def forward(self, x):\n        B, N, D = x.shape\n        qkv = self.to_qkv(x).view(B, N, 3, self.num_heads, self.head_dim).permute(2,0,3,1,4)\n        Q, K, V = qkv[0], qkv[1], qkv[2]  # each (B,h,N,d)\n        attn_out = attention(Q, K, V)  # (B,h,N,d)\n        attn_out = attn_out.permute(0,2,1,3).contiguous().view(B, N, D)\n        return self.to_out(attn_out)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class FeedForward(nn.Module):\n    def __init__(self, dim, hidden_dim=2048):\n        super().__init__()\n        self.net = nn.Sequential(\n            nn.Linear(dim, hidden_dim),\n            nn.GELU(),\n            nn.Linear(hidden_dim, dim)\n        )\n    def forward(self, x):\n        return self.net(x)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class PerceiverBlock(nn.Module):\n    def __init__(self, latent_dim, input_dim, num_heads, chunk_size):\n        super().__init__()\n        self.ln_cross = nn.LayerNorm(latent_dim)\n        self.cross_attn = CrossAttention(latent_dim, input_dim, num_heads, chunk_size)\n        self.ln_self = nn.LayerNorm(latent_dim)\n        self.self_attn = SelfAttention(latent_dim, num_heads)\n        self.ln_ffn = nn.LayerNorm(latent_dim)\n        self.ffn = FeedForward(latent_dim)\n        \n    def forward(self, latents, byte_arr):\n        latents = latents + self.cross_attn(self.ln_cross(latents), byte_arr)\n        latents = latents + self.self_attn(self.ln_self(latents))\n        latents = latents + self.ffn(self.ln_ffn(latents))\n        return latents","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class Perceiver(nn.Module):\n    def __init__(self, num_classes=1000, input_chans=3, num_latents=512, latent_dim=512, num_blocks=8,\n                 num_heads=8, chunk_size=2048, num_bands=64):\n        super().__init__()\n        self.encode = FourierEncode(num_bands=num_bands)\n        self.latent_init = nn.Parameter(torch.randn(1, num_latents, latent_dim))\n        self.num_blocks = num_blocks\n        input_dim = input_chans + 4 * num_bands\n        self.perceiver_block = PerceiverBlock(latent_dim=latent_dim, input_dim=input_dim, num_heads=num_heads, chunk_size=chunk_size)\n        self.norm = nn.LayerNorm(latent_dim)\n        self.classifier = nn.Linear(latent_dim, num_classes)\n\n    def forward(self, pixels, coords):\n        B, M, _ = pixels.shape\n        pos_enc = self.encode(coords)\n        byte_arr = torch.cat([pixels, pos_enc], dim=-1)\n        latents = self.latent_init.expand(B, -1, -1)\n\n        for _ in range(self.num_blocks):\n            latents = self.perceiver_block(latents, byte_arr)\n\n        z = latents.mean(dim=1)\n        z = self.norm(z)\n        logits = self.classifier(z)\n        return logits","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch_xla\nimport torch_xla.core.xla_model as xm\nimport torch_xla.distributed.parallel_loader as pl\nimport torch_xla.amp\nimport torch_xla.runtime as xr","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"criterion = nn.CrossEntropyLoss(label_smoothing=config.label_smoothing)\n\ndef get_lr(epoch, optim_steps_per_epoch, opt_step):\n    global_step = epoch * optim_steps_per_epoch + opt_step\n    warmup_steps = config.warmup_epochs * optim_steps_per_epoch\n    total_steps = config.num_epochs * optim_steps_per_epoch\n    if global_step < warmup_steps:\n        return config.base_lr * float(global_step) / float(max(1, warmup_steps))\n    else:\n        progress = float(global_step - warmup_steps) / float(max(1, total_steps - warmup_steps))\n        return config.base_lr * 0.5 * (1.0 + np.cos(np.pi * progress))","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_loader, val_loader, num_classes = build_dataloaders(config)\n# Wrap loaders for TPU\ndevice = torch_xla.device()\ntrain_device_loader = pl.MpDeviceLoader(train_loader, device)\nval_device_loader = pl.MpDeviceLoader(val_loader, device)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model = Perceiver()\nmodel = model.to(device)\n\noptimizer = optim.AdamW(model.parameters(), lr=config.base_lr, betas=(config.adam_beta1, config.adam_beta2), weight_decay=config.weight_decay)\n\ndef train_epoch(model, loader, optimizer, criterion, epoch, step_per_epoch):\n    model.train()\n    running_loss = 0.0\n    total_samples = 0\n    optimizer.zero_grad()\n    accumulation_steps = config.gradient_accumulation_steps\n    optim_steps_per_epoch = len(loader) // accumulation_steps\n    opt_step = 0\n    \n    for step, (pixels, coords, labels) in enumerate(loader):\n        pixels, coords, labels = pixels.to(device), coords.to(device), labels.to(device)\n        global_step = epoch * step_per_epoch + step\n\n        with torch.autocast(\"xla\", dtype=torch.bfloat16):\n            logits = model(pixels, coords)\n            loss = criterion(logits, labels) / accumulation_steps\n\n        loss.backward()\n        running_loss += loss.item() * pixels.size(0) * accumulation_steps\n        total_samples += labels.size(0)\n        \n        if (step + 1) % accumulation_steps == 0:\n            lr = get_lr(epoch, optim_steps_per_epoch, opt_step)\n            for param_group in optimizer.param_groups:\n                param_group['lr'] = lr\n            xm.optimizer_step(optimizer, barrier=True)\n            optimizer.zero_grad()\n            opt_step += 1\n            if xm.is_master_ordinal():\n                print(f\"Epoch [{epoch+1}/{config.num_epochs}] Step [{step+1}/{step_per_epoch}]\" \n                f\"Loss: {running_loss/((step+1)*config.batch_size_per_core*xr.world_size()):.4f}\")\n\n        if step % 20 == 0:\n            xm.master_print(f'Epoch {epoch+1}, Step {step}, Loss: {loss.item():.4f}')\n\n        if step % 100 == 0:\n            torch_xla.sync()\n\n    avg_loss = running_loss / total_samples\n    return avg_loss","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def validate(model, loader, criterion, epoch):\n    model.eval()\n    correct_top1 = 0\n    correct_top5 = 0\n    total = 0\n    val_loss = 0.0\n    with torch.no_grad():\n        for pixels, coords, labels in loader:\n            with torch.autocast(\"xla\", dtype=torch.bfloat16):\n                logits = model(pixels, coords)\n                loss = criterion(logits, labels)\n                \n            val_loss += loss.item() * pixels.size(0)\n            _, preds = torch.topk(logits, k=5, dim=-1)\n            correct_top1 += preds[:, 0].eq(labels).sum().item()\n\n            for i in range(labels.size(0)):\n                if labels[i] in preds[i]:\n                    correct_top5 += 1\n            total += labels.size(0)\n            \n    avg_loss = val_loss / total\n    top1 = correct_top1 / total\n    top5 = correct_top5 / total\n    \n    xm.master_print(f\"Validation Epoch {epoch+1}: Loss {avg_loss:.4f}, Top1 {top1:.4f}, Top5 {top5:.4f}\")\n\n    return top1, top5","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"best_top1 = 0.0\n\ntotal_batch_size = config.batch_size_per_core * xr.world_size()\nconfig.gradient_accumulation_steps = config.effective_batch_size // total_batch_size\n\nstep_per_epoch = len(train_loader) // config.gradient_accumulation_steps\n\nfor epoch in range(config.num_epochs):\n    train_loss = train_epoch(model, train_device_loader, optimizer, criterion, epoch, step_per_epoch)\n    top1, top5 = validate(model, val_device_loader, criterion, epoch)\n    if top1 > best_top1:\n        best_top1 = top1","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}