{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":14774,"databundleVersionId":875431}],"dockerImageVersionId":31287,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-03-09T15:58:19.384862Z","iopub.execute_input":"2026-03-09T15:58:19.385718Z","iopub.status.idle":"2026-03-09T15:58:27.155307Z","shell.execute_reply.started":"2026-03-09T15:58:19.385684Z","shell.execute_reply":"2026-03-09T15:58:27.154440Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nprint(os.listdir(\"/kaggle/input/competitions\"))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-09T15:58:27.156838Z","iopub.execute_input":"2026-03-09T15:58:27.157107Z","iopub.status.idle":"2026-03-09T15:58:27.161899Z","shell.execute_reply.started":"2026-03-09T15:58:27.157081Z","shell.execute_reply":"2026-03-09T15:58:27.161240Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(os.listdir(\"/kaggle/input/competitions/aptos2019-blindness-detection\"))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-09T15:58:27.162909Z","iopub.execute_input":"2026-03-09T15:58:27.163152Z","iopub.status.idle":"2026-03-09T15:58:27.174606Z","shell.execute_reply.started":"2026-03-09T15:58:27.163130Z","shell.execute_reply":"2026-03-09T15:58:27.173865Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================\n# CELL 1 — INSTALL (optional) + IMPORTS\n# =========================\n\n# Kaggle biasanya sudah punya torch, pandas, sklearn, PIL\n# Uncomment kalau butuh package tambahan\n# !pip -q install torch torchvision\n\nimport os\nimport copy\nimport math\nimport random\nimport numpy as np\nimport pandas as pd\nfrom PIL import Image\nfrom collections import OrderedDict, defaultdict\n\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import accuracy_score, classification_report, confusion_matrix\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.autograd import Function\nfrom torch.utils.data import Dataset, DataLoader\n\nimport torchvision\nfrom torchvision import transforms, models\n\nprint(\"Torch version:\", torch.__version__)\nprint(\"Torchvision version:\", torchvision.__version__)\nprint(\"CUDA available:\", torch.cuda.is_available())\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(\"Using device:\", device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-09T15:58:27.175552Z","iopub.execute_input":"2026-03-09T15:58:27.175775Z","iopub.status.idle":"2026-03-09T15:58:27.187973Z","shell.execute_reply.started":"2026-03-09T15:58:27.175735Z","shell.execute_reply":"2026-03-09T15:58:27.187129Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================\n# CELL 2 — CONFIG\n# =========================\n\nSEED = 42\nrandom.seed(SEED)\nnp.random.seed(SEED)\ntorch.manual_seed(SEED)\ntorch.cuda.manual_seed_all(SEED)\n\nNPZ_PATH = \"/kaggle/input/orion-aptos-npz/orion_dr_224.npz\"\n\nIMG_SIZE = 224\nBATCH_SIZE = 16\nNUM_WORKERS = 2\n\nNUM_CLASSES = 5\nNUM_CLIENTS = 3\nCLIENT_RATIOS = [1, 2, 3]          # pembagian data client\nCLIENT_SEQUENCE_MODE = \"ascending\" # pilihan: \"ascending\", \"descending\", \"fixed\"\n\nROUNDS = 3\nLOCAL_EPOCHS = 1\nLR = 1e-4\nWEIGHT_DECAY = 1e-4\n\nDANN_LAMBDA = 0.2\nVAL_SIZE = 0.2\n\nprint(\"NPZ_PATH:\", NPZ_PATH)\nprint(\"NUM_CLIENTS:\", NUM_CLIENTS)\nprint(\"CLIENT_RATIOS:\", CLIENT_RATIOS)\nprint(\"CLIENT_SEQUENCE_MODE:\", CLIENT_SEQUENCE_MODE)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-09T15:58:27.189860Z","iopub.execute_input":"2026-03-09T15:58:27.190141Z","iopub.status.idle":"2026-03-09T15:58:27.207610Z","shell.execute_reply.started":"2026-03-09T15:58:27.190118Z","shell.execute_reply":"2026-03-09T15:58:27.206758Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================\n# CELL 2.1 — BUILD NPZ FROM APTOS DATASET\n# =========================\n\nimport os\nimport numpy as np\nimport pandas as pd\nfrom PIL import Image\nfrom tqdm import tqdm\n\nDATA_ROOT = \"/kaggle/input/competitions/aptos2019-blindness-detection\"\n\nCSV_PATH = os.path.join(DATA_ROOT, \"train.csv\")\nIMG_DIR  = os.path.join(DATA_ROOT, \"train_images\")\n\nIMG_SIZE = 224\n\nSAVE_NPZ = \"/kaggle/working/orion_dr_224.npz\"\n\nprint(\"Loading CSV...\")\ndf = pd.read_csv(CSV_PATH)\n\nprint(\"Total samples:\", len(df))\n\nimages = []\nlabels = []\nmissing = 0\n\nfor _, row in tqdm(df.iterrows(), total=len(df)):\n\n    img_path = os.path.join(IMG_DIR, row[\"id_code\"] + \".png\")\n\n    if not os.path.exists(img_path):\n        missing += 1\n        continue\n\n    try:\n        img = Image.open(img_path).convert(\"RGB\")\n        img = img.resize((IMG_SIZE, IMG_SIZE))\n\n        img = np.array(img)\n\n        images.append(img)\n        labels.append(row[\"diagnosis\"])\n\n    except:\n        missing += 1\n\nimages = np.array(images)\nlabels = np.array(labels)\n\nprint(\"\\nImages shape:\", images.shape)\nprint(\"Labels shape:\", labels.shape)\nprint(\"Missing images:\", missing)\n\nprint(\"\\nSaving NPZ...\")\n\nnp.savez_compressed(\n    SAVE_NPZ,\n    images=images,\n    labels=labels\n)\n\nprint(\"Saved to:\", SAVE_NPZ)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-09T15:58:27.208739Z","iopub.execute_input":"2026-03-09T15:58:27.209075Z","iopub.status.idle":"2026-03-09T16:08:50.859002Z","shell.execute_reply.started":"2026-03-09T15:58:27.209037Z","shell.execute_reply":"2026-03-09T16:08:50.857844Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================\n# CELL 2.2 — VERIFY NPZ\n# =========================\n\nimport numpy as np\n\nNPZ_PATH = \"/kaggle/working/orion_dr_224.npz\"\n\ndata = np.load(NPZ_PATH)\n\nprint(\"Keys:\", data.files)\n\nimages = data[\"images\"]\nlabels = data[\"labels\"]\n\nprint(\"Images shape:\", images.shape)\nprint(\"Labels shape:\", labels.shape)\n\nprint(\"Class distribution:\")\nprint(pd.Series(labels).value_counts().sort_index())\n\n# =========================\n# CELL 3 — LOAD NPZ\n# =========================\n\nimport numpy as np\n\nNPZ_PATH = \"/kaggle/working/orion_dr_224.npz\"\n\ndata = np.load(NPZ_PATH)\n\nimages = data[\"images\"]\nlabels = data[\"labels\"]\n\nprint(\"Images shape:\", images.shape)\nprint(\"Labels shape:\", labels.shape)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-09T16:08:50.860318Z","iopub.execute_input":"2026-03-09T16:08:50.860627Z","iopub.status.idle":"2026-03-09T16:08:59.578098Z","shell.execute_reply.started":"2026-03-09T16:08:50.860599Z","shell.execute_reply":"2026-03-09T16:08:59.577165Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================\n# CELL 3 — SANITY CHECK NPZ\n# =========================\n\nprint(\"NPZ exists:\", os.path.exists(NPZ_PATH))\n\ndata = np.load(NPZ_PATH, allow_pickle=True)\nprint(\"Keys in NPZ:\", data.files)\n\nfor k in data.files:\n    print(k, data[k].shape if hasattr(data[k], \"shape\") else type(data[k]))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-09T16:08:59.579302Z","iopub.execute_input":"2026-03-09T16:08:59.579982Z","iopub.status.idle":"2026-03-09T16:09:08.241657Z","shell.execute_reply.started":"2026-03-09T16:08:59.579946Z","shell.execute_reply":"2026-03-09T16:09:08.240784Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================\n# CELL 4 — LOAD NPZ\n# =========================\n\ndata = np.load(NPZ_PATH, allow_pickle=True)\n\n# sesuaikan dengan key di NPZ kamu\nimages = data[\"images\"]   # atau data[\"X\"]\nlabels = data[\"labels\"]   # atau data[\"y\"]\n\nprint(\"images shape:\", images.shape)\nprint(\"labels shape:\", labels.shape)\nprint(\"label distribution:\", pd.Series(labels).value_counts().sort_index())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-09T16:09:08.242658Z","iopub.execute_input":"2026-03-09T16:09:08.242915Z","iopub.status.idle":"2026-03-09T16:09:12.577515Z","shell.execute_reply.started":"2026-03-09T16:09:08.242890Z","shell.execute_reply":"2026-03-09T16:09:12.576468Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================\n# CELL 5 — TRAIN / VAL SPLIT\n# =========================\n\nindices = np.arange(len(labels))\n\ntrain_idx, val_idx = train_test_split(\n    indices,\n    test_size=VAL_SIZE,\n    random_state=SEED,\n    stratify=labels\n)\n\nX_train = images[train_idx]\ny_train = labels[train_idx]\n\nX_val = images[val_idx]\ny_val = labels[val_idx]\n\nprint(\"Train:\", X_train.shape, y_train.shape)\nprint(\"Val  :\", X_val.shape, y_val.shape)\n\nprint(\"\\nTrain class dist:\")\nprint(pd.Series(y_train).value_counts().sort_index())\n\nprint(\"\\nVal class dist:\")\nprint(pd.Series(y_val).value_counts().sort_index())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-09T16:09:12.578751Z","iopub.execute_input":"2026-03-09T16:09:12.579057Z","iopub.status.idle":"2026-03-09T16:09:12.720411Z","shell.execute_reply.started":"2026-03-09T16:09:12.579030Z","shell.execute_reply":"2026-03-09T16:09:12.719581Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================\n# CELL 6 — SPLIT TRAIN DATA INTO CLIENTS WITH RATIO 1:2:3\n# =========================\n\ndef make_client_splits_by_ratio_stratified(X, y, client_ratios, seed=42):\n    rng = np.random.default_rng(seed)\n\n    num_clients = len(client_ratios)\n    ratio_sum = sum(client_ratios)\n\n    client_indices = [[] for _ in range(num_clients)]\n\n    for cls in np.unique(y):\n        cls_idx = np.where(y == cls)[0]\n        rng.shuffle(cls_idx)\n\n        n_cls = len(cls_idx)\n\n        raw_counts = np.array(client_ratios, dtype=float) / ratio_sum * n_cls\n        counts = np.floor(raw_counts).astype(int)\n\n        remainder = n_cls - counts.sum()\n        if remainder > 0:\n            frac_parts = raw_counts - counts\n            add_order = np.argsort(-frac_parts)\n            for i in range(remainder):\n                counts[add_order[i]] += 1\n\n        start = 0\n        for cid in range(num_clients):\n            end = start + counts[cid]\n            client_indices[cid].extend(cls_idx[start:end].tolist())\n            start = end\n\n    client_data = []\n    for cid in range(num_clients):\n        idx = np.array(client_indices[cid])\n        rng.shuffle(idx)\n\n        X_c = X[idx]\n        y_c = y[idx]\n        d_c = np.full(len(idx), cid, dtype=np.int64)\n\n        client_data.append({\n            \"images\": X_c,\n            \"labels\": y_c,\n            \"domains\": d_c\n        })\n\n    return client_data\n\nclient_data = make_client_splits_by_ratio_stratified(\n    X_train,\n    y_train,\n    client_ratios=CLIENT_RATIOS,\n    seed=SEED\n)\n\nfor cid, cdata in enumerate(client_data):\n    print(f\"\\nClient {cid}: {len(cdata['labels'])} samples\")\n    print(\"Class distribution:\")\n    print(pd.Series(cdata[\"labels\"]).value_counts().sort_index())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-09T16:09:12.721812Z","iopub.execute_input":"2026-03-09T16:09:12.722496Z","iopub.status.idle":"2026-03-09T16:09:13.084212Z","shell.execute_reply.started":"2026-03-09T16:09:12.722440Z","shell.execute_reply":"2026-03-09T16:09:13.083476Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================\n# CELL 7 — DATASET FOR NPZ\n# =========================\n\ntrain_tfms = transforms.Compose([\n    transforms.ToPILImage(),\n    transforms.Resize((IMG_SIZE, IMG_SIZE)),\n    transforms.RandomHorizontalFlip(p=0.5),\n    transforms.RandomRotation(degrees=10),\n    transforms.ColorJitter(brightness=0.1, contrast=0.1),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=[0.485, 0.456, 0.406],\n                         std =[0.229, 0.224, 0.225]),\n])\n\nval_tfms = transforms.Compose([\n    transforms.ToPILImage(),\n    transforms.Resize((IMG_SIZE, IMG_SIZE)),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=[0.485, 0.456, 0.406],\n                         std =[0.229, 0.224, 0.225]),\n])\n\n\nclass AptosNPZDANNDataset(Dataset):\n    def __init__(self, images, labels, domains=None, transform=None):\n        self.images = images\n        self.labels = labels\n        self.domains = domains if domains is not None else np.zeros(len(labels), dtype=np.int64)\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.labels)\n\n    def __getitem__(self, idx):\n        img = self.images[idx]\n\n        # pastikan uint8 kalau memang hasil preprocess berupa pixel image\n        if img.dtype != np.uint8:\n            if img.max() <= 1.0:\n                img = (img * 255).astype(np.uint8)\n            else:\n                img = img.astype(np.uint8)\n\n        if self.transform is not None:\n            img = self.transform(img)\n\n        label = int(self.labels[idx])\n        domain = int(self.domains[idx])\n\n        return {\n            \"image\": img,\n            \"label\": torch.tensor(label, dtype=torch.long),\n            \"domain\": torch.tensor(domain, dtype=torch.long)\n        }","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-09T16:09:13.085352Z","iopub.execute_input":"2026-03-09T16:09:13.085796Z","iopub.status.idle":"2026-03-09T16:09:13.095640Z","shell.execute_reply.started":"2026-03-09T16:09:13.085745Z","shell.execute_reply":"2026-03-09T16:09:13.094745Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================\n# CELL 8 — BUILD DATALOADERS\n# =========================\n\ndef build_loader_np(images, labels, domains=None, transform=None, batch_size=BATCH_SIZE, shuffle=False):\n    ds = AptosNPZDANNDataset(\n        images=images,\n        labels=labels,\n        domains=domains,\n        transform=transform\n    )\n    dl = DataLoader(\n        ds,\n        batch_size=batch_size,\n        shuffle=shuffle,\n        num_workers=NUM_WORKERS,\n        pin_memory=torch.cuda.is_available(),\n        drop_last=False\n    )\n    return ds, dl\n\nclient_loaders = {}\nclient_sizes = {}\n\nfor cid, cdata in enumerate(client_data):\n    ds, dl = build_loader_np(\n        images=cdata[\"images\"],\n        labels=cdata[\"labels\"],\n        domains=cdata[\"domains\"],\n        transform=train_tfms,\n        batch_size=BATCH_SIZE,\n        shuffle=True\n    )\n    client_loaders[cid] = dl\n    client_sizes[cid] = len(ds)\n\nval_domains = np.zeros(len(y_val), dtype=np.int64)\nval_ds, val_loader = build_loader_np(\n    images=X_val,\n    labels=y_val,\n    domains=val_domains,\n    transform=val_tfms,\n    batch_size=BATCH_SIZE,\n    shuffle=False\n)\n\nprint(\"Validation size:\", len(val_ds))\nprint(\"Client sizes:\", client_sizes)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-09T16:09:13.096701Z","iopub.execute_input":"2026-03-09T16:09:13.097002Z","iopub.status.idle":"2026-03-09T16:09:13.115388Z","shell.execute_reply.started":"2026-03-09T16:09:13.096977Z","shell.execute_reply":"2026-03-09T16:09:13.114641Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================\n# CELL 9 — GRADIENT REVERSAL LAYER\n# =========================\n\nclass GradientReversalFunction(Function):\n    @staticmethod\n    def forward(ctx, x, alpha):\n        ctx.alpha = alpha\n        return x.view_as(x)\n\n    @staticmethod\n    def backward(ctx, grad_output):\n        return grad_output.neg() * ctx.alpha, None\n\n\nclass GRL(nn.Module):\n    def forward(self, x, alpha=1.0):\n        return GradientReversalFunction.apply(x, alpha)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-09T16:09:13.118854Z","iopub.execute_input":"2026-03-09T16:09:13.119242Z","iopub.status.idle":"2026-03-09T16:09:13.129122Z","shell.execute_reply.started":"2026-03-09T16:09:13.119209Z","shell.execute_reply":"2026-03-09T16:09:13.128403Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================\n# CELL 10 — DANN MODEL\n# =========================\n\nclass FeatureExtractor(nn.Module):\n    def __init__(self, out_dim=256):\n        super().__init__()\n        backbone = models.resnet18(weights=models.ResNet18_Weights.DEFAULT)\n        in_features = backbone.fc.in_features\n        backbone.fc = nn.Identity()\n        self.backbone = backbone\n        self.proj = nn.Sequential(\n            nn.Linear(in_features, out_dim),\n            nn.ReLU(),\n            nn.Dropout(0.2)\n        )\n\n    def forward(self, x):\n        feat = self.backbone(x)\n        feat = self.proj(feat)\n        return feat\n\n\nclass LabelClassifier(nn.Module):\n    def __init__(self, in_dim=256, num_classes=5):\n        super().__init__()\n        self.net = nn.Sequential(\n            nn.Linear(in_dim, 128),\n            nn.ReLU(),\n            nn.Dropout(0.2),\n            nn.Linear(128, num_classes)\n        )\n\n    def forward(self, x):\n        return self.net(x)\n\n\nclass DomainClassifier(nn.Module):\n    def __init__(self, in_dim=256, num_domains=5):\n        super().__init__()\n        self.net = nn.Sequential(\n            nn.Linear(in_dim, 128),\n            nn.ReLU(),\n            nn.Dropout(0.2),\n            nn.Linear(128, num_domains)\n        )\n\n    def forward(self, x):\n        return self.net(x)\n\n\nclass DANNModel(nn.Module):\n    def __init__(self, feat_dim=256, num_classes=5, num_domains=5):\n        super().__init__()\n        self.feature_extractor = FeatureExtractor(out_dim=feat_dim)\n        self.label_classifier = LabelClassifier(in_dim=feat_dim, num_classes=num_classes)\n        self.domain_classifier = DomainClassifier(in_dim=feat_dim, num_domains=num_domains)\n        self.grl = GRL()\n\n    def forward(self, x, alpha=1.0):\n        feat = self.feature_extractor(x)\n        class_logits = self.label_classifier(feat)\n        rev_feat = self.grl(feat, alpha)\n        domain_logits = self.domain_classifier(rev_feat)\n        return class_logits, domain_logits","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-09T16:09:13.130030Z","iopub.execute_input":"2026-03-09T16:09:13.130296Z","iopub.status.idle":"2026-03-09T16:09:13.143182Z","shell.execute_reply.started":"2026-03-09T16:09:13.130258Z","shell.execute_reply":"2026-03-09T16:09:13.142414Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================\n# CELL 11 — FEDERATED UTILS\n# =========================\n\ndef get_model_state(model):\n    return copy.deepcopy(model.state_dict())\n\ndef set_model_state(model, state_dict):\n    model.load_state_dict(copy.deepcopy(state_dict))\n\ndef fedavg(state_dicts, weights):\n    avg_state = OrderedDict()\n    total_weight = sum(weights)\n\n    for k in state_dicts[0].keys():\n        avg_state[k] = sum(sd[k] * (w / total_weight) for sd, w in zip(state_dicts, weights))\n    return avg_state","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-09T16:09:13.144055Z","iopub.execute_input":"2026-03-09T16:09:13.144375Z","iopub.status.idle":"2026-03-09T16:09:13.161201Z","shell.execute_reply.started":"2026-03-09T16:09:13.144351Z","shell.execute_reply":"2026-03-09T16:09:13.160348Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================\n# NEW CELL 11.5 — CLIENT SEQUENCING HELPER\n# =========================\n\ndef get_client_sequence(client_sizes, mode=\"ascending\"):\n    valid_clients = [cid for cid, size in client_sizes.items() if size > 0]\n\n    if mode == \"ascending\":\n        ordered = sorted(valid_clients, key=lambda cid: client_sizes[cid])\n    elif mode == \"descending\":\n        ordered = sorted(valid_clients, key=lambda cid: client_sizes[cid], reverse=True)\n    elif mode == \"fixed\":\n        ordered = valid_clients\n    else:\n        raise ValueError(f\"Unknown CLIENT_SEQUENCE_MODE: {mode}\")\n\n    return ordered\n\nclient_order = get_client_sequence(client_sizes, mode=CLIENT_SEQUENCE_MODE)\n\nprint(\"Client sizes:\", client_sizes)\nprint(\"Client order:\", client_order)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-09T16:09:13.162205Z","iopub.execute_input":"2026-03-09T16:09:13.162553Z","iopub.status.idle":"2026-03-09T16:09:13.176613Z","shell.execute_reply.started":"2026-03-09T16:09:13.162494Z","shell.execute_reply":"2026-03-09T16:09:13.175768Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================\n# CELL 12 — LOCAL TRAINING (DANN)\n# =========================\n\ndef train_one_client_dann(global_state, client_loader, client_id, num_domains=NUM_CLIENTS,\n                          local_epochs=1, lr=1e-4, weight_decay=1e-4, dann_lambda=0.2):\n    model = DANNModel(\n        feat_dim=256,\n        num_classes=NUM_CLASSES,\n        num_domains=num_domains\n    ).to(device)\n\n    set_model_state(model, global_state)\n\n    optimizer = torch.optim.Adam(model.parameters(), lr=lr, weight_decay=weight_decay)\n    criterion_cls = nn.CrossEntropyLoss()\n    criterion_dom = nn.CrossEntropyLoss()\n\n    model.train()\n\n    total_loss = 0.0\n    total_cls_loss = 0.0\n    total_dom_loss = 0.0\n    total_correct = 0\n    total_count = 0\n\n    for epoch in range(local_epochs):\n        for batch_idx, batch in enumerate(client_loader):\n            images = batch[\"image\"].to(device)\n            labels = batch[\"label\"].to(device)\n            domains = batch[\"domain\"].to(device)\n\n            # progress alpha buat GRL\n            p = float(batch_idx + epoch * len(client_loader)) / max(1, (local_epochs * len(client_loader)))\n            alpha = 2. / (1. + np.exp(-10 * p)) - 1\n\n            optimizer.zero_grad()\n\n            class_logits, domain_logits = model(images, alpha=alpha)\n\n            cls_loss = criterion_cls(class_logits, labels)\n            dom_loss = criterion_dom(domain_logits, domains)\n            loss = cls_loss + dann_lambda * dom_loss\n\n            loss.backward()\n            optimizer.step()\n\n            total_loss += loss.item() * images.size(0)\n            total_cls_loss += cls_loss.item() * images.size(0)\n            total_dom_loss += dom_loss.item() * images.size(0)\n\n            preds = class_logits.argmax(dim=1)\n            total_correct += (preds == labels).sum().item()\n            total_count += labels.size(0)\n\n    metrics = {\n        \"loss\": total_loss / max(1, total_count),\n        \"cls_loss\": total_cls_loss / max(1, total_count),\n        \"dom_loss\": total_dom_loss / max(1, total_count),\n        \"acc\": total_correct / max(1, total_count)\n    }\n\n    return get_model_state(model), metrics","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-09T16:09:13.177375Z","iopub.execute_input":"2026-03-09T16:09:13.177647Z","iopub.status.idle":"2026-03-09T16:09:13.189106Z","shell.execute_reply.started":"2026-03-09T16:09:13.177623Z","shell.execute_reply":"2026-03-09T16:09:13.188382Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================\n# CELL 13 — EVALUATION\n# =========================\n\n@torch.no_grad()\ndef evaluate_global_model(state_dict, data_loader):\n    model = DANNModel(\n        feat_dim=256,\n        num_classes=NUM_CLASSES,\n        num_domains=NUM_CLIENTS\n    ).to(device)\n    set_model_state(model, state_dict)\n\n    model.eval()\n\n    all_preds = []\n    all_labels = []\n    total_loss = 0.0\n    total_count = 0\n\n    criterion = nn.CrossEntropyLoss()\n\n    for batch in data_loader:\n        images = batch[\"image\"].to(device)\n        labels = batch[\"label\"].to(device)\n\n        class_logits, _ = model(images, alpha=0.0)\n        loss = criterion(class_logits, labels)\n\n        preds = class_logits.argmax(dim=1)\n\n        total_loss += loss.item() * images.size(0)\n        total_count += labels.size(0)\n\n        all_preds.extend(preds.cpu().numpy().tolist())\n        all_labels.extend(labels.cpu().numpy().tolist())\n\n    acc = accuracy_score(all_labels, all_preds)\n\n    return {\n        \"loss\": total_loss / max(1, total_count),\n        \"acc\": acc,\n        \"y_true\": all_labels,\n        \"y_pred\": all_preds\n    }","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-09T16:09:13.190028Z","iopub.execute_input":"2026-03-09T16:09:13.190326Z","iopub.status.idle":"2026-03-09T16:09:13.206923Z","shell.execute_reply.started":"2026-03-09T16:09:13.190297Z","shell.execute_reply":"2026-03-09T16:09:13.206041Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================\n# CELL 14 — INIT GLOBAL MODEL\n# =========================\n\nglobal_model = DANNModel(\n    feat_dim=256,\n    num_classes=NUM_CLASSES,\n    num_domains=NUM_CLIENTS\n).to(device)\n\nglobal_state = get_model_state(global_model)\n\nprint(\"Global model initialized.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-09T16:09:13.208096Z","iopub.execute_input":"2026-03-09T16:09:13.208743Z","iopub.status.idle":"2026-03-09T16:09:14.141873Z","shell.execute_reply.started":"2026-03-09T16:09:13.208708Z","shell.execute_reply":"2026-03-09T16:09:14.141052Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================\n# CELL 15 — FEDERATED TRAINING LOOP WITH CLIENT SEQUENCING\n# =========================\n\nhistory = []\n\nfor rnd in range(1, ROUNDS + 1):\n    print(f\"\\n{'='*50}\")\n    print(f\"ROUND {rnd}/{ROUNDS}\")\n    print(f\"{'='*50}\")\n\n    client_order = get_client_sequence(client_sizes, mode=CLIENT_SEQUENCE_MODE)\n    print(\"Sequence this round:\", client_order)\n\n    local_states = []\n    local_weights = []\n    round_logs = []\n\n    for cid in client_order:\n        if client_sizes[cid] == 0:\n            continue\n\n        local_state, metrics = train_one_client_dann(\n            global_state=global_state,\n            client_loader=client_loaders[cid],\n            client_id=cid,\n            num_domains=NUM_CLIENTS,\n            local_epochs=LOCAL_EPOCHS,\n            lr=LR,\n            weight_decay=WEIGHT_DECAY,\n            dann_lambda=DANN_LAMBDA\n        )\n\n        local_states.append(local_state)\n        local_weights.append(client_sizes[cid])\n        round_logs.append((cid, metrics))\n\n        print(\n            f\"Client {cid} | \"\n            f\"size={client_sizes[cid]} | \"\n            f\"loss={metrics['loss']:.4f} | \"\n            f\"cls_loss={metrics['cls_loss']:.4f} | \"\n            f\"dom_loss={metrics['dom_loss']:.4f} | \"\n            f\"acc={metrics['acc']:.4f}\"\n        )\n\n    global_state = fedavg(local_states, local_weights)\n\n    val_metrics = evaluate_global_model(global_state, val_loader)\n\n    print(f\"\\n[GLOBAL VAL] loss={val_metrics['loss']:.4f} | acc={val_metrics['acc']:.4f}\")\n\n    history.append({\n        \"round\": rnd,\n        \"client_order\": client_order,\n        \"client_logs\": round_logs,\n        \"val_loss\": val_metrics[\"loss\"],\n        \"val_acc\": val_metrics[\"acc\"]\n    })","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-09T16:09:14.142858Z","iopub.execute_input":"2026-03-09T16:09:14.143212Z","iopub.status.idle":"2026-03-09T16:09:50.742822Z","shell.execute_reply.started":"2026-03-09T16:09:14.143186Z","shell.execute_reply":"2026-03-09T16:09:50.741674Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================\n# CELL 16 — FINAL REPORT\n# =========================\n\nfinal_metrics = evaluate_global_model(global_state, val_loader)\n\nprint(\"Final Validation Loss:\", final_metrics[\"loss\"])\nprint(\"Final Validation Accuracy:\", final_metrics[\"acc\"])\n\nprint(\"\\nClassification Report:\")\nprint(classification_report(final_metrics[\"y_true\"], final_metrics[\"y_pred\"], digits=4))\n\nprint(\"\\nConfusion Matrix:\")\nprint(confusion_matrix(final_metrics[\"y_true\"], final_metrics[\"y_pred\"]))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-09T16:09:50.744339Z","iopub.execute_input":"2026-03-09T16:09:50.744815Z","iopub.status.idle":"2026-03-09T16:09:51.993654Z","shell.execute_reply.started":"2026-03-09T16:09:50.744775Z","shell.execute_reply":"2026-03-09T16:09:51.992626Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================\n# CELL 17 — SAVE MODEL\n# =========================\n\nSAVE_PATH = \"/kaggle/working/fl_dann_aptos_resnet18.pth\"\ntorch.save(global_state, SAVE_PATH)\nprint(\"Saved to:\", SAVE_PATH)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-09T16:09:51.995463Z","iopub.execute_input":"2026-03-09T16:09:51.995934Z","iopub.status.idle":"2026-03-09T16:09:52.072508Z","shell.execute_reply.started":"2026-03-09T16:09:51.995898Z","shell.execute_reply":"2026-03-09T16:09:52.071720Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================\n# CELL 18 — SAVE HISTORY\n# =========================\n\nrows = []\nfor h in history:\n    rows.append({\n        \"round\": h[\"round\"],\n        \"val_loss\": h[\"val_loss\"],\n        \"val_acc\": h[\"val_acc\"]\n    })\n\nhist_df = pd.DataFrame(rows)\ndisplay(hist_df)\n\nhist_df.to_csv(\"/kaggle/working/fl_dann_history.csv\", index=False)\nprint(\"History saved.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-09T16:09:52.073807Z","iopub.execute_input":"2026-03-09T16:09:52.074148Z","iopub.status.idle":"2026-03-09T16:09:52.107596Z","shell.execute_reply.started":"2026-03-09T16:09:52.074111Z","shell.execute_reply":"2026-03-09T16:09:52.106858Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}