{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":100385,"databundleVersionId":12076007,"sourceType":"competition"}],"dockerImageVersionId":31012,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"Code sample for FID computation on hyperspectral data, ","metadata":{}},{"cell_type":"markdown","source":"To adapt the Fréchet Inception Distance (FID) to hyperspectral data, we first transform each 125-band cube (spanning 450–950 nm) into a three-band image that Inception V3 can accept. Rather than discarding most of the spectrum by naively choosing three arbitrary bands, every pixel’s spectrum is convolved with the official Sentinel-2 spectral-response functions for bands B3 (Green, centred ≈ 560 nm), B4 (Red, ≈ 665 nm) and B8 (NIR, ≈ 842 nm). Because the SRF weights are normalised, this operation is an energy-preserving integral that faithfully simulates the radiance those Sentinel-2 detectors would record, compressing all 125 wavelengths into physically interpretable Bule, Green, Red and NIR channels.\n\nWe choose Green, Red, and NIR bands to organise a false-color 3-band image. The resulting false-color image is resized to 299 × 299 px and normalised with ImageNet means and standard deviations before being passed through Inception V3 up to the global-average-pool layer, yielding a 2048-dimensional embedding for each sample. For the real and generated hyperspectral data we compute the mean vector and covariance of these embeddings and plug them into the standard FID formula, which measures the Fréchet distance between the two multivariate Gaussians. This pipeline preserves the full spectral information, links the bands to well-understood remote-sensing physics, and still leverages the widely used Inception feature space, giving a single quantitative score that reflects how closely the synthetic cubes resemble real hyperspectral data in both spectral and spatial structure.\n\nThe code can be find below:","metadata":{}},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torchvision.models import inception_v3, Inception_V3_Weights\nfrom torchvision import transforms\nimport numpy as np\nfrom scipy.linalg import sqrtm\n\nSRF_GREEN = torch.tensor([\n    0.0000,0.0000,0.0000,0.0000,0.0001,0.0002,0.0005,0.0008,0.0014,0.0024,0.0041,\n    0.0069,0.0113,0.0180,0.0279,0.0414,0.0583,0.0783,0.1008,0.1252,0.1507,0.1766,\n    0.2023,0.2271,0.2505,0.2721,0.2913,0.3079,0.3216,0.3324,0.3404,0.3459,0.3495,\n    0.3516,0.3528,0.3533,0.3535,0.3536,0.3538,0.3539,0.3541,0.3542,0.3542,0.3541,\n    0.3535,0.3520,0.3491,0.3443,0.3373,0.3277,0.3152,0.2997,0.2811,0.2595,0.2349,\n    0.2076,0.1778,0.1462,0.1140,0.0823,0.0524,0.0259,0.0037,0.0003,0.0000,0.0000,\n    0.0000,0.0000,0.0000,0.0000,0.0000,0.0000,0.0000,0.0000,0.0000,0.0000,\n]).float()\n\nSRF_RED = torch.tensor([\n    0.0000,0.0000,0.0000,0.0000,0.0001,0.0002,0.0003,0.0006,0.0012,0.0024,0.0047,\n    0.0087,0.0154,0.0255,0.0395,0.0575,0.0786,0.1020,0.1265,0.1505,0.1732,0.1940,\n    0.2121,0.2269,0.2381,0.2454,0.2491,0.2494,0.2466,0.2409,0.2326,0.2219,0.2093,\n    0.1952,0.1799,0.1639,0.1476,0.1314,0.1157,0.1008,0.0870,0.0744,0.0629,0.0525,\n    0.0430,0.0344,0.0266,0.0195,0.0129,0.0070,0.0018,0.0003,0.0000,0.0000,0.0000,\n    0.0000,0.0000,0.0000,0.0000,0.0000,0.0000,0.0000,0.0000,0.0000,0.0000,\n]).float()\n\nSRF_NIR = torch.tensor([\n    0.0000,0.0000,0.0000,0.0000,0.0000,0.0001,0.0002,0.0003,0.0006,0.0011,0.0022,\n    0.0041,0.0073,0.0125,0.0204,0.0317,0.0470,0.0666,0.0905,0.1185,0.1500,0.1841,\n    0.2196,0.2554,0.2900,0.3219,0.3495,0.3715,0.3870,0.3950,0.3950,0.3872,0.3721,\n    0.3503,0.3228,0.2912,0.2573,0.2228,0.1888,0.1563,0.1261,0.0990,0.0755,0.0557,\n    0.0395,0.0265,0.0162,0.0082,0.0023,0.0003,0.0000,0.0000,0.0000,0.0000,0.0000,\n    0.0000,0.0000,0.0000,0.0000,0.0000,0.0000,0.0000,0.0000,0.0000,0.0000,\n]).float()\n\nSRF_TABLE = {\n    'green': SRF_GREEN / SRF_GREEN.sum(),\n    'red'  : SRF_RED   / SRF_RED.sum(),\n    'nir'  : SRF_NIR   / SRF_NIR.sum(),\n}\n\nWAVELENGTHS = torch.linspace(450.0, 950.0, 125)  # [125]\n\ndef resample_srf(srf_1d: torch.Tensor,\n                 wl_axis: torch.Tensor) -> torch.Tensor:\n    wl_start, wl_stop = 450.0, 950.0\n    N = srf_1d.numel()\n    xp = np.linspace(wl_start, wl_stop, N)\n    fp = srf_1d.cpu().numpy()\n    values = np.interp(wl_axis.cpu().numpy(), xp, fp)\n    out = torch.from_numpy(values).float()\n    return out / out.sum()                              # re-normalise\n\nSRF_RESAMPLED = {k: resample_srf(v, WAVELENGTHS) for k, v in SRF_TABLE.items()}\n\ndef hs_to_s2_rgb(hs_img: torch.Tensor) -> torch.Tensor:\n\n    if hs_img.shape[0] != 125:\n        raise ValueError('Expect 125 spectral bands (450-950 nm).')\n    H, W = hs_img.shape[1:]\n    out = []\n    for key in ('green', 'red', 'nir'):\n        w = SRF_RESAMPLED[key].view(125, 1, 1).to(hs_img.device)\n        out.append((hs_img * w).sum(0))\n    return torch.stack(out)\n\nclass InceptionPool3(nn.Module):\n\n    def __init__(self, device):\n        super().__init__()\n\n        # load pre-trained weights (aux_logits **must** stay True)\n        weights = Inception_V3_Weights.IMAGENET1K_V1\n        net = inception_v3(weights=weights,\n                           aux_logits=True,\n                           transform_input=False).to(device)\n        net.eval()\n\n        net.AuxLogits = nn.Identity()\n\n        self.stem_and_blocks = nn.Sequential(*list(net.children())[:-2])\n        self.avgpool         = net.avgpool\n        self.norm = transforms.Normalize(mean=[0.485, 0.456, 0.406],\n                                         std =[0.229, 0.224, 0.225])\n\n    @torch.no_grad()\n    def forward(self, x):\n        x = self.norm(x)\n        x = self.stem_and_blocks(x)\n        x = self.avgpool(x)\n        return torch.flatten(x, 1)\n\n\ndef get_activations(hs_list, model, device='cpu', batch=8):\n    feats = []\n    with torch.no_grad():\n        for i in range(0, len(hs_list), batch):\n            hs = torch.stack(hs_list[i:i+batch]).to(device)\n            rgb = torch.stack([hs_to_s2_rgb(img) for img in hs])\n            rgb = F.interpolate(rgb, size=(299, 299), mode='bilinear', align_corners=False)\n            feats.append(model(rgb).cpu().numpy())\n    return np.concatenate(feats, axis=0)\n\ndef stats(a): return a.mean(0), np.cov(a, rowvar=False)\n\ndef fid(mu1, sigma1, mu2, sigma2):\n    diff = mu1 - mu2\n    covmean, _ = sqrtm(sigma1 @ sigma2, disp=False)\n    if np.iscomplexobj(covmean):\n        covmean = covmean.real\n    return diff @ diff + np.trace(sigma1 + sigma2 - 2.0 * covmean)\n\n\nif __name__ == '__main__':\n    device = 'cuda' if torch.cuda.is_available() else 'cpu'\n\n    gen_images = samples\n    real_images = image\n\n    inc = InceptionPool3(device)\n\n    act_real = get_activations(real_images, inc, device)\n    act_gen  = get_activations(gen_images,  inc, device)\n\n    mu_r, sig_r = stats(act_real)\n    mu_g, sig_g = stats(act_gen)\n\n    print(f'FID : {fid(mu_r, sig_r, mu_g, sig_g):.4f}')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":" The above example will calculate one FID value from given gen_images (50 hyperspectral patches for example) and real_images (also 50 as we provided) corresponding to a given disease level. Attendants will calulate all of the 10 FIDs corresponding to 10 disease levels, and submit the results.","metadata":{}},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}