{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":5048,"databundleVersionId":868335,"isSourceIdPinned":false,"sourceType":"competition"}],"dockerImageVersionId":30674,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"from transformers import ViTImageProcessor, ViTForImageClassification ,ViTConfig\nfrom torch.utils.data import DataLoader\nfrom torch.optim import AdamW\nfrom torch.optim.lr_scheduler import StepLR\nimport torch\nfrom tqdm import tqdm\nimport torch.nn.functional as F\nfrom torch.utils.data import DataLoader, ConcatDataset\nimport os\nimport pandas as pd\nfrom sklearn.model_selection import train_test_split\nfrom torchvision import transforms\nfrom torch.utils.data import Dataset, DataLoader\nfrom PIL import Image\nimport copy\ndevice = \"cuda\" if torch.cuda.is_available() else \"cpu\"\ndevice","metadata":{"_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","execution":{"iopub.status.busy":"2025-10-05T12:19:24.144194Z","iopub.execute_input":"2025-10-05T12:19:24.144541Z","iopub.status.idle":"2025-10-05T12:19:30.325258Z","shell.execute_reply.started":"2025-10-05T12:19:24.144511Z","shell.execute_reply":"2025-10-05T12:19:30.324381Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Specify the path to your CSV file\nfile_path = '/kaggle/input/state-farm-distracted-driver-detection/driver_imgs_list.csv'\nprocessor = ViTImageProcessor.from_pretrained('google/vit-base-patch16-224')\n\n# Read the CSV file\ndff = pd.read_csv(file_path)\ndistinct_values = dff['subject'].unique()\ndistinct_count = len(distinct_values)\n# Get the number of rows\nnum_rows = len(dff)\n\nprint(f\"The number of rows in the CSV file is: {num_rows}\")\nprint(f\"clients ID: {distinct_values}\")\nprint(f\"Number of clients : {distinct_count}\")\nprint(dff.columns)\nsubject_df = dff[dff['subject'] == distinct_values[1]]\nprint(subject_df)","metadata":{"execution":{"iopub.status.busy":"2025-10-05T12:19:30.326788Z","iopub.execute_input":"2025-10-05T12:19:30.327289Z","iopub.status.idle":"2025-10-05T12:19:30.49224Z","shell.execute_reply.started":"2025-10-05T12:19:30.327263Z","shell.execute_reply":"2025-10-05T12:19:30.49145Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"clients_df={}\n\nclass_labels = {\n    0: \"safe driving\",\n    1: \"texting - right\",\n    2: \"talking on the phone - right\",\n    3: \"texting - left\",\n    4: \"talking on the phone - left\",\n    5: \"operating the radio\",\n    6: \"drinking\",\n    7: \"reaching behind\",\n    8: \"hair and makeup\",\n    9: \"talking to passenger\"\n}\nTraining_data='/kaggle/input/state-farm-distracted-driver-detection/imgs/train'\n\n\nfor idx, subject in enumerate(distinct_values, start=1):\n    subject_df = dff[dff['subject'] == subject]\n    \n    data_dict = {\n    'image_path': [],\n    'class': [],\n    }\n    \n    for _, row in subject_df.iterrows():\n        class_path = os.path.join(Training_data, row['classname'])\n        image_path = os.path.join(class_path, row['img'])\n        classname = row['classname'].strip()\n        number_part = int(classname[1:])\n        data_dict['image_path'].append(image_path)\n        data_dict['class'].append(number_part)\n        \n    df_client = pd.DataFrame(data_dict)\n    clients_df[f\"df_client{idx}\"] = df_client","metadata":{"execution":{"iopub.status.busy":"2025-10-05T12:19:30.493115Z","iopub.execute_input":"2025-10-05T12:19:30.493345Z","iopub.status.idle":"2025-10-05T12:19:31.663225Z","shell.execute_reply.started":"2025-10-05T12:19:30.493324Z","shell.execute_reply":"2025-10-05T12:19:31.662467Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(clients_df[\"df_client1\"])\nprint(clients_df[\"df_client2\"])","metadata":{"execution":{"iopub.status.busy":"2025-10-05T12:19:31.664935Z","iopub.execute_input":"2025-10-05T12:19:31.66522Z","iopub.status.idle":"2025-10-05T12:19:31.67328Z","shell.execute_reply.started":"2025-10-05T12:19:31.665197Z","shell.execute_reply":"2025-10-05T12:19:31.672343Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"summ=0\nfor name, df in clients_df.items():\n    print(f\"{name} has {len(df)} rows\")\n    summ=summ+len(df)\nprint(f\"total number = {summ}\")","metadata":{"execution":{"iopub.status.busy":"2025-10-05T12:19:31.674243Z","iopub.execute_input":"2025-10-05T12:19:31.674525Z","iopub.status.idle":"2025-10-05T12:19:31.68538Z","shell.execute_reply.started":"2025-10-05T12:19:31.674505Z","shell.execute_reply":"2025-10-05T12:19:31.684455Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class StateFarmDataset(Dataset):\n    def __init__(self, dataframe, processor):\n        self.data = dataframe\n        self.processor = processor\n\n    def __len__(self):\n        return len(self.data)\n\n    def __getitem__(self, index):\n        row = self.data.iloc[index]\n        image_path = row['image_path']\n        label = row['class']\n\n        # Load and process the image\n        image = Image.open(image_path).convert(\"RGB\")\n        inputs = self.processor(images=image, return_tensors=\"pt\")\n        return inputs['pixel_values'].squeeze(0), label","metadata":{"execution":{"iopub.status.busy":"2025-10-05T12:19:31.686461Z","iopub.execute_input":"2025-10-05T12:19:31.686703Z","iopub.status.idle":"2025-10-05T12:19:31.697659Z","shell.execute_reply.started":"2025-10-05T12:19:31.686685Z","shell.execute_reply":"2025-10-05T12:19:31.696666Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.model_selection import train_test_split\nfrom torch.utils.data import DataLoader\n\n# Lists to store dataloaders\nclients_traindataloader = []\nclients_valdataloader = []\nclients_testdataloader = []\n# Iterate over the clients_df dictionary\nfor key, df in clients_df.items():\n    # Split the DataFrame into train, validation, and test\n    train_df, val_df = train_test_split(df, test_size=0.3, random_state=42, shuffle=True)\n    vall_df, test_df= train_test_split(val_df, test_size=0.333, random_state=42, shuffle=True)\n\n    train_dataset = StateFarmDataset(train_df, processor)\n    train_dataloader = DataLoader(dataset=train_dataset, batch_size=16, shuffle=True, num_workers=4)\n    \n    val_dataset = StateFarmDataset(vall_df, processor)\n    val_dataloader = DataLoader(dataset=val_dataset, batch_size=16, shuffle=False, num_workers=4)\n\n    test_dataset = StateFarmDataset(test_df, processor)\n    test_dataloader = DataLoader(dataset=test_dataset, batch_size=16, shuffle=False, num_workers=4)\n    \n    clients_traindataloader.append(train_dataloader)\n    clients_valdataloader.append(val_dataloader)\n    clients_testdataloader.append(test_dataloader)","metadata":{"execution":{"iopub.status.busy":"2025-10-05T12:19:31.698701Z","iopub.execute_input":"2025-10-05T12:19:31.698941Z","iopub.status.idle":"2025-10-05T12:19:31.750784Z","shell.execute_reply.started":"2025-10-05T12:19:31.698923Z","shell.execute_reply":"2025-10-05T12:19:31.749955Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for idx, dataloader in enumerate(clients_traindataloader):\n    print(f\"Client {idx} - Train samples: {len(dataloader.dataset)}\")\n\nfor idx, dataloader in enumerate(clients_valdataloader):\n    print(f\"Client {idx} - Validation samples: {len(dataloader.dataset)}\")\n\nfor idx, dataloader in enumerate(clients_testdataloader):\n    print(f\"Client {idx} - Test samples: {len(dataloader.dataset)}\")\n","metadata":{"execution":{"iopub.status.busy":"2025-10-05T12:19:31.751787Z","iopub.execute_input":"2025-10-05T12:19:31.752017Z","iopub.status.idle":"2025-10-05T12:19:31.757255Z","shell.execute_reply.started":"2025-10-05T12:19:31.751997Z","shell.execute_reply":"2025-10-05T12:19:31.756392Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport numpy as np\n\n# Function to show images\ndef imshow(img, label):\n    # Convert the tensor to a numpy array and rescale to [0, 1]\n    img = img / 2 + 0.5  # Unnormalize\n    np_img = img.numpy()\n\n    return np.transpose(np_img, (1, 2, 0))  # Convert the image from (C, H, W) to (H, W, C)\n\n# Number of images to show per client\nnum_images = 5\n\nclient1_data = next(iter(clients_traindataloader[0]))\nclient2_data = next(iter(clients_traindataloader[1]))\nclient3_data = next(iter(clients_traindataloader[3]))\n\n# Unpack the data into images and labels\nimages1, labels1 = client1_data\nimages2, labels2 = client2_data\nimages3, labels3 = client3_data\n# Create a figure with two rows: one for client 1 and one for client 2\nfig, axes = plt.subplots(3, num_images, figsize=(15, 6))\n\n# Plot images from client 1\nfor i in range(num_images):\n    label1 = labels1[i].item()\n    axes[0, i].imshow(imshow(images1[i], labels1[i]))\n    axes[0, i].set_title(f'  Client 1: Class {class_labels[label1]}')\n    axes[0, i].axis('off')  # Hide axes\n\n# Plot images from client 2\nfor i in range(num_images):\n    label2 = labels2[i].item()\n    axes[1, i].imshow(imshow(images2[i], labels2[i]))\n    axes[1, i].set_title(f'  Client 2: Class {class_labels[label2]}')\n    axes[1, i].axis('off')  # Hide axes\n    \nfor i in range(num_images):\n    label3 = labels3[i].item()\n    axes[2, i].imshow(imshow(images3[i], labels3[i]))\n    axes[2, i].set_title(f'  Client 3: Class {class_labels[label3]}')\n    axes[2, i].axis('off')\n    \n# Adjust layout for better presentation\nplt.tight_layout()\nplt.show()\n","metadata":{"execution":{"iopub.status.busy":"2025-10-05T12:19:31.758495Z","iopub.execute_input":"2025-10-05T12:19:31.759033Z","iopub.status.idle":"2025-10-05T12:19:35.083207Z","shell.execute_reply.started":"2025-10-05T12:19:31.759003Z","shell.execute_reply":"2025-10-05T12:19:35.082266Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"****Histogram For validation of every client model on each data of all clients to ensure the highest accuracy on it's data****","metadata":{}},{"cell_type":"code","source":"import matplotlib.pyplot as plt\n\ndef Histogram_validation(model, clients_valdataloader, when, device):\n    accuracies = []  # List to store accuracy for each client\n    client_indices = []  # List to store client indices\n    \n    for idx, val_dataloader in enumerate(clients_valdataloader):\n        model.eval()  # Set model to evaluation mode\n        total_loss = 0.0\n        correct_predictions = 0\n        total_samples = 0  # Track total number of samples\n\n        with torch.no_grad():  # No need to track gradients for validation\n            for images, labels in val_dataloader:\n                images = images.to(device)\n                labels = labels.to(device)\n\n                # Forward pass\n                outputs = model(images)\n                logits = outputs.logits\n\n                # Calculate the loss\n                loss = F.cross_entropy(logits, labels)\n                total_loss += loss.item()\n              \n                predicted_class_idx = logits.argmax(-1)\n                correct_predictions += (predicted_class_idx == labels).sum().item()\n                total_samples += labels.size(0)  # Increment total_samples by batch size\n\n        val_accuracy = correct_predictions / total_samples * 100  # Accuracy in percentage\n        accuracies.append(val_accuracy)  # Append accuracy for the current client\n        client_indices.append(idx + 1)  # Store the client index (1-based)\n        print(f\"Client {idx + 1} Validation Loss {when}: {total_loss / len(val_dataloader):.4f}, Validation Accuracy {when}: {val_accuracy:.2f}%\")\n    \n    # Plotting the accuracy with client index on the x-axis\n    plt.figure(figsize=(8, 6))\n    plt.bar(client_indices, accuracies, color='skyblue', edgecolor='black')\n    plt.title(f'Validation Accuracy by Client {when}')\n    plt.xlabel('Client Index')\n    plt.ylabel('Accuracy (%)')\n    plt.xticks(client_indices)  # Make sure all client indices appear as ticks on the x-axis\n    plt.grid(True, axis='y')\n    plt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-05T12:19:35.086124Z","iopub.execute_input":"2025-10-05T12:19:35.086399Z","iopub.status.idle":"2025-10-05T12:19:35.095154Z","shell.execute_reply.started":"2025-10-05T12:19:35.086372Z","shell.execute_reply":"2025-10-05T12:19:35.094254Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch.nn.functional as F\nimport torch\n\ndef validate(model, val_dataloader, device):\n    model.eval()\n    total_loss = 0.0\n    correct_predictions = 0\n    total_samples = 0\n    all_probs = []  # Store probability distributions\n    all_labels = []  # Store true labels for comparison\n\n    with torch.no_grad():\n        for images, labels in val_dataloader:\n            images, labels = images.to(device), labels.to(device)\n\n            # Forward pass\n            outputs = model(images)\n            logits = outputs.logits\n\n            # Compute softmax probabilities\n            probs = F.softmax(logits, dim=1)\n            all_probs.append(probs.cpu())  # Store probabilities in CPU memory\n            all_labels.append(labels.cpu())\n\n            # Calculate loss\n            loss = F.cross_entropy(logits, labels)\n            total_loss += loss.item()\n            \n            # Compute accuracy\n            predicted_class_idx = logits.argmax(-1)\n            correct_predictions += (predicted_class_idx == labels).sum().item()\n            total_samples += labels.size(0)\n\n    # Convert lists to tensors\n    all_probs = torch.cat(all_probs, dim=0)  # Shape: (num_samples, num_classes)\n    all_labels = torch.cat(all_labels, dim=0)  # Shape: (num_samples,)\n\n    # Compute validation accuracy and loss\n    val_accuracy = correct_predictions / total_samples * 100\n    val_loss = total_loss / len(val_dataloader)\n\n    return val_loss, val_accuracy, all_probs, all_labels\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-05T12:19:35.096195Z","iopub.execute_input":"2025-10-05T12:19:35.096457Z","iopub.status.idle":"2025-10-05T12:19:35.108201Z","shell.execute_reply.started":"2025-10-05T12:19:35.096423Z","shell.execute_reply":"2025-10-05T12:19:35.107542Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def Server(global_weights, client_weights, client_sizes, clients_loss):\n    lora_keys_server = [key for key in global_weights.keys() if \"lora\" in key or \"classifier\"in key]\n    # print(f\"lora keys in server {lora_keys_server}\")\n    total_loss=sum(clients_loss)\n    total_client_size=sum(client_sizes)\n    # Compute global weights for all clients\n    for key in lora_keys_server:\n        global_weights[key] = sum(\n            (client_weights[i][key] * (clients_loss[i]*1 / total_loss))\n            for i in range(len(client_weights))\n        )\n    return global_weights","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-05T12:19:35.109068Z","iopub.execute_input":"2025-10-05T12:19:35.109261Z","iopub.status.idle":"2025-10-05T12:19:35.132228Z","shell.execute_reply.started":"2025-10-05T12:19:35.109245Z","shell.execute_reply":"2025-10-05T12:19:35.13154Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch.nn.functional as F\n\ndef Server_kl(global_weights, client_weights, client_sizes, clients_loss, clients_prob, clients_labels, global_probs, global_labels):\n    lora_keys_server_KL = [key for key in global_weights.keys() if \"lora\" in key or \"classifier\" in key]\n    # print(f\"lora keys in KL {lora_keys_server_KL}\")\n    total_loss = sum(clients_loss)\n    total_client_size = sum(client_sizes)\n\n    # Compute KL divergence for each client\n    kl_divergences = []\n    for i in range(len(client_weights)):\n        if global_probs[i].shape != clients_prob[i].shape:\n            raise ValueError(f\"Shape mismatch: global_probs[{i}] has shape {global_probs[i].shape}, but clients_prob[{i}] has shape {clients_prob[i].shape}\")\n\n        kl_div = F.kl_div(\n            global_probs[i].log(),  # Log-prob of client\n            clients_prob[i],  # Global model prob\n            reduction='batchmean'\n        )\n        kl_divergences.append(kl_div.item())\n\n    total_kl = sum(kl_divergences)\n\n    # Aggregate global weights using loss and KL divergence\n    for key in lora_keys_server_KL:\n        global_weights[key] = sum(\n            (client_weights[i][key] * ((kl_divergences[i]*1.07 / total_kl)))\n            for i in range(len(client_weights))\n        )\n\n    return global_weights,kl_divergences\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-05T12:19:35.132995Z","iopub.execute_input":"2025-10-05T12:19:35.133196Z","iopub.status.idle":"2025-10-05T12:19:35.145132Z","shell.execute_reply.started":"2025-10-05T12:19:35.13318Z","shell.execute_reply":"2025-10-05T12:19:35.144466Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport copy\nimport torch\nimport torch.nn.functional as F\nfrom sklearn.cluster import KMeans\nfrom torch.utils.data import ConcatDataset, DataLoader\n\nOUTPUT_DIR = \"/kaggle/working/models\"\nos.makedirs(OUTPUT_DIR, exist_ok=True)\n\ndef save_client_model(client_idx, model):\n    torch.save(model.state_dict(), f\"{OUTPUT_DIR}/client_{client_idx}.pt\")\n\ndef load_client_model(client_idx, model, device):\n    model.load_state_dict(torch.load(f\"{OUTPUT_DIR}/client_{client_idx}.pt\", map_location=device))\n    model.to(device)\n    return model\n\ndef compute_similarity_matrix(client_weights_cpu):\n    vectors = []\n    for weights in client_weights_cpu:\n        flat = torch.cat([v.flatten() for v in weights.values()])\n        vectors.append(flat)\n    return torch.stack(vectors)\n\n\ndef train(client_models, globalmodel, clients_traindataloader, clients_valdataloader,\n                             client_optimizers, client_scheduler, num_epochs, device, k=3):\n    \n    validation_results = {i: {\"All_data\": [], \"after\": []} for i in range(len(client_models))}\n    Best_Clients_Model = {i: {\"model\": globalmodel, \"loss\": 30} for i in range(len(client_models))}\n\n    for epoch in range(num_epochs):\n        print(f\"\\nEpoch {epoch + 1}/{num_epochs}\")\n        client_weights, client_sizes = [], []\n\n        # === 1. Train and Save ===\n        for idx, traindataloader in enumerate(clients_traindataloader):\n            model = client_models[idx]\n            model.train()\n            optimizer = client_optimizers[idx]\n            total_loss, correct, total = 0.0, 0, 0\n\n            for images, labels in traindataloader:\n                images, labels = images.to(device), labels.to(device)\n                optimizer.zero_grad()\n                logits = model(images).logits\n                loss = F.cross_entropy(logits, labels)\n                loss.backward()\n                optimizer.step()\n                total_loss += loss.item()\n                correct += (logits.argmax(dim=1) == labels).sum().item()\n                total += labels.size(0)\n\n            if client_scheduler[idx]:\n                client_scheduler[idx].step()\n\n            acc = correct / total * 100\n            print(f\"Client {idx} Training Loss={total_loss:.4f}, Acc={acc:.2f}%\")\n\n            lora_params = {k: v.cpu() for k, v in model.state_dict().items() if \"lora\" in k or \"classifier\" in k}\n            client_weights.append(copy.deepcopy(lora_params))\n            client_sizes.append(total)\n\n            save_client_model(idx, model)\n            del model\n            torch.cuda.empty_cache()\n\n        # === 2. Cluster & Aggregate ===\n        if (epoch + 1) % 2 == 0 or (epoch + 1) == num_epochs:\n            matrix = compute_similarity_matrix(client_weights)\n            kmeans = KMeans(n_clusters=k, random_state=42)\n            cluster_labels = kmeans.fit_predict(matrix)\n            cluster_models = {}\n\n            for cluster_id in range(k):\n                cluster_indices = [i for i, label in enumerate(cluster_labels) if label == cluster_id]\n                if not cluster_indices:\n                    continue\n\n                print(f\"\\n-- Cluster {cluster_id}: {cluster_indices}\")\n                c_weights = [client_weights[i] for i in cluster_indices]\n                c_sizes = [client_sizes[i] for i in cluster_indices]\n                c_loss = []\n                c_probs, c_labels = [], []\n\n                # Step 1: Validate before KL aggregation\n                for i in cluster_indices:\n                    model = load_client_model(i, client_models[i], device)\n                    val_loss, val_acc, probs, labels = validate(model, clients_valdataloader[i], device)\n                    print(f\"Client {i} Validation BEFORE KL: Loss={val_loss:.4f}, Acc={val_acc:.2f}%\")\n                    c_loss.append(val_loss)\n                    c_probs.append(probs)\n                    c_labels.append(labels)\n\n                    save_client_model(i, model)\n                    del model\n                    torch.cuda.empty_cache()\n\n                # Step 2: First aggregation (Server)\n                agg_weights = Server(globalmodel.state_dict(), c_weights, c_sizes, c_loss)\n                cluster_models[cluster_id] = agg_weights\n\n                # # Step 3: Load model weights and validate cluster global on concatenated data\n                # for i in cluster_indices:\n                #     model = client_models[i]\n                #     state_dict = model.state_dict()\n                #     for key in agg_weights:\n                #         state_dict[key] = agg_weights[key]\n                #     model.load_state_dict(state_dict)\n                #     model.to(device)\n                # Step 3: Load aggregated weights into each client in cluster and validate on its own data\n                cluster_client_probs = []\n                cluster_client_labels = []\n                \n                for i in cluster_indices:\n                    model = client_models[i]\n                    state_dict = model.state_dict()\n                    for key in agg_weights:\n                        state_dict[key] = agg_weights[key]\n                    model.load_state_dict(state_dict)\n                    model.to(device)\n                \n                    val_loss, val_acc, probs, labels = validate(model, clients_valdataloader[i], device)\n                    print(f\"Client {i} Validation AFTER Server aggregation: Loss={val_loss:.4f}, Acc={val_acc:.2f}%\")\n                \n                    cluster_client_probs.append(probs)\n                    cluster_client_labels.append(labels)\n\n        \n                # # Step 4: KL-based aggregation\n                # agg_weights_kl, _ = Server_kl(agg_weights, c_weights, c_sizes, c_loss, c_probs, c_labels,\n                #                               [g_probs]*len(cluster_indices), [g_labels]*len(cluster_indices))\n                agg_weights_kl, _ = Server_kl(\n                    agg_weights,             # aggregated weights\n                    c_weights,               # pre-agg client weights\n                    c_sizes,                 # data sizes\n                    c_loss,                  # validation loss before aggregation\n                    c_probs,                 # client predictions before agg\n                    c_labels,                # client labels before agg\n                    cluster_client_probs,    # client predictions after Server agg\n                    cluster_client_labels    # client labels after Server agg\n                )\n\n                cluster_models[cluster_id] = agg_weights_kl  # update\n\n\n                # Step 5: Load updated model and final validate on clients\n                for i in cluster_indices:\n                    model = client_models[i]\n                    state_dict = model.state_dict()\n                    for key in agg_weights_kl:\n                        state_dict[key] = agg_weights_kl[key]\n                    model.load_state_dict(state_dict)\n                    model.to(device)\n\n                    val_loss, val_acc, _, _ = validate(model, clients_valdataloader[i], device)\n                    print(f\"Client {i} Validation AFTER KL: Loss={val_loss:.4f}, Acc={val_acc:.2f}%\")\n                    validation_results[i][\"after\"].append((val_loss, val_acc))\n\n                    if val_loss < Best_Clients_Model[i][\"loss\"]:\n                        Best_Clients_Model[i][\"model\"] = copy.deepcopy({\n                            k: v.cpu() for k, v in model.state_dict().items()\n                        })\n                        Best_Clients_Model[i][\"loss\"] = val_loss\n\n                # # Final validation of each cluster model on its cluster data\n                # val_loss, val_acc, _, _ = validate(client_models[cluster_indices[0]], concat_loader, device)\n                # print(f\"Cluster {cluster_id} Global Model AFTER KL on CONCAT: Loss={val_loss:.4f}, Acc={val_acc:.2f}%\")\n                # validation_results[cluster_indices[0]][\"All_data\"].append((val_loss, val_acc))\n                # Final validation of each cluster model on the FULL concatenated data\n                full_dataset = ConcatDataset([dl.dataset for dl in clients_valdataloader])\n                full_loader = DataLoader(full_dataset, batch_size=16, shuffle=False)\n                \n                val_loss, val_acc, _, _ = validate(client_models[cluster_indices[0]], full_loader, device)\n                print(f\"Cluster {cluster_id} Global Model on FULL CONCAT DATA: Loss={val_loss:.4f}, Acc={val_acc:.2f}%\")\n                \n\n            # Step 6: Rollback if not final epoch\n            if (epoch + 1) != num_epochs:\n                for i in range(len(client_models)):\n                    client_models[i].load_state_dict(Best_Clients_Model[i][\"model\"])\n                    client_models[i].to(device)\n\n    return validation_results, cluster_models, Best_Clients_Model , cluster_labels\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-05T12:19:35.146459Z","iopub.execute_input":"2025-10-05T12:19:35.147041Z","iopub.status.idle":"2025-10-05T12:19:35.220451Z","shell.execute_reply.started":"2025-10-05T12:19:35.147011Z","shell.execute_reply":"2025-10-05T12:19:35.219559Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"****without lora paramaters= 85.806.346****\n                                                \n                                                \n****LORA paramaters=2666656****","metadata":{}},{"cell_type":"code","source":"# architecture of LORA\nimport math\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport numpy as np\nfrom transformers import ViTForImageClassification\nfrom torch.optim import AdamW\nfrom torch.optim.lr_scheduler import StepLR\nimport random\nimport warnings\nwarnings.filterwarnings('ignore')\n\n# ----------------------------------------------------------\n# LoRA Components\n# ----------------------------------------------------------\nclass LoRALayer():\n    def __init__(self, r: int, lora_alpha: int, lora_dropout: float, merge_weights: bool):\n        self.r = r\n        self.lora_alpha = lora_alpha\n        if lora_dropout > 0.:\n            self.lora_dropout = nn.Dropout(p=lora_dropout)\n        else:\n            self.lora_dropout = lambda x: x\n        self.merged = False\n        self.merge_weights = merge_weights\n\n\nclass Linear(nn.Linear, LoRALayer):\n    def __init__(self, in_features: int, out_features: int, r: int = 0, lora_alpha: int = 1,\n                 lora_dropout: float = 0., fan_in_fan_out: bool = False, merge_weights: bool = True, **kwargs):\n        nn.Linear.__init__(self, in_features, out_features, **kwargs)\n        LoRALayer.__init__(self, r=r, lora_alpha=lora_alpha, lora_dropout=lora_dropout, merge_weights=merge_weights)\n        self.fan_in_fan_out = fan_in_fan_out\n        if r > 0:\n            self.lora_B = nn.Parameter(self.weight.new_zeros((out_features, r)))\n            self.lora_A = nn.Parameter(torch.randn(r, in_features))\n            self.scaling = self.lora_alpha / self.r\n            self.weight.requires_grad = False\n        self.reset_parameters()\n        if fan_in_fan_out:\n            self.weight.data = self.weight.data.transpose(0, 1)\n\n    def reset_parameters(self):\n        nn.Linear.reset_parameters(self)\n        if hasattr(self, 'lora_A'):\n            nn.init.kaiming_uniform_(self.lora_A, a=math.sqrt(5))\n            nn.init.zeros_(self.lora_B)\n\n    def forward(self, x: torch.Tensor):\n        def T(w):\n            return w.transpose(0, 1) if self.fan_in_fan_out else w\n        if self.r > 0 and not self.merged:\n            result = F.linear(x, T(self.weight), bias=self.bias)\n            result += (self.lora_dropout(x) @ self.lora_A.T @ self.lora_B.T) * self.scaling\n            return result\n        else:\n            return F.linear(x, T(self.weight), bias=self.bias)\n\n\n# ----------------------------------------------------------\n# LoRA Vision Transformer\n# ----------------------------------------------------------\nclass LoRAViT(ViTForImageClassification):\n    def __init__(self, *args, **kwargs):\n        super(LoRAViT, self).__init__(*args, **kwargs)\n\n        # ✅ Reduce ViT encoder layers from 12 to 6\n        self.vit.encoder.layer = self.vit.encoder.layer[:6]\n\n        # LoRA classifier\n        self.classifier = Linear(self.config.hidden_size, 10, r=32, lora_alpha=16)\n        self.config.num_labels = 10\n\n        # Replace Linear layers with LoRA versions\n        linear_layers = []\n        for name, module in self.named_modules():\n            if isinstance(module, nn.Linear):\n                linear_layers.append((name, module))\n\n        i = 0\n        for name, module in linear_layers:\n            in_features, out_features = module.in_features, module.out_features\n            fan_in_fan_out = getattr(module, 'fan_in_fan_out', False)\n            if i < 3:\n                dropout = 0\n            elif i < 6:\n                dropout = 0.15\n            elif i < 11:\n                dropout = 0.2\n            else:\n                dropout = 0\n            lora_layer = Linear(in_features, out_features, r=32, lora_dropout=dropout,\n                                lora_alpha=16, fan_in_fan_out=fan_in_fan_out)\n            i += 1\n\n            parent_module = self\n            name_parts = name.split(\".\")\n            for part in name_parts[:-1]:\n                parent_module = getattr(parent_module, part)\n\n            lora_layer.weight.data = module.weight.data.clone()\n            if module.bias is not None:\n                lora_layer.bias.data = module.bias.data.clone()\n            setattr(parent_module, name_parts[-1], lora_layer)\n\n\n# ----------------------------------------------------------\n# Reproducibility\n# ----------------------------------------------------------\nSEED = 42\nrandom.seed(SEED)\nnp.random.seed(SEED)\ntorch.manual_seed(SEED)\ntorch.cuda.manual_seed(SEED)\ntorch.backends.cudnn.deterministic = True\ntorch.backends.cudnn.benchmark = False\n\n# ----------------------------------------------------------\n# Model Initialization\n# ----------------------------------------------------------\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\nbase_model = ViTForImageClassification.from_pretrained(\"google/vit-base-patch16-224\")\nbase_model.vit.encoder.layer = base_model.vit.encoder.layer[:6]  # ✅ truncate base model as well\npretrained_classifier_weight = base_model.classifier.weight[:10].clone()\npretrained_classifier_bias = base_model.classifier.bias[:10].clone()\n\nglobalmodel = LoRAViT.from_pretrained('google/vit-base-patch16-224', ignore_mismatched_sizes=True).to(device)\nglobalmodel.classifier.weight.data.copy_(pretrained_classifier_weight)\nglobalmodel.classifier.bias.data.copy_(pretrained_classifier_bias)\n\nclient_models = [LoRAViT.from_pretrained('google/vit-base-patch16-224', ignore_mismatched_sizes=True).to(device)\n                 for _ in range(26)]\nfor model in client_models:\n    model.classifier.weight.data.copy_(pretrained_classifier_weight)\n    model.classifier.bias.data.copy_(pretrained_classifier_bias)\n\n# ----------------------------------------------------------\n# Optimizers\n# ----------------------------------------------------------\nclient_optimizers = []\nclient_scheduler = []\n\nfor model in client_models:\n    loRA_params = []\n    for name, param in model.named_parameters():\n        if \"lora\" in name or \"classifier\" in name:\n            param.requires_grad = True\n            loRA_params.append(param)\n        else:\n            param.requires_grad = False\n    optimizer = AdamW(loRA_params, lr=1e-3)\n    client_optimizers.append(optimizer)\n    scheduler = StepLR(optimizer, step_size=7, gamma=0.7)\n    client_scheduler.append(scheduler)\n\n# ----------------------------------------------------------\n# Sanity Check\n# ----------------------------------------------------------\nprint(\"✅ Encoder layers count:\", len(globalmodel.vit.encoder.layer))  # should print 6\ntotal_trainable = sum(p.numel() for p in globalmodel.parameters() if p.requires_grad)\nprint(f\"Total trainable parameters: {total_trainable}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-05T12:19:35.221728Z","iopub.execute_input":"2025-10-05T12:19:35.221944Z","iopub.status.idle":"2025-10-05T12:19:55.306337Z","shell.execute_reply.started":"2025-10-05T12:19:35.221926Z","shell.execute_reply":"2025-10-05T12:19:55.305436Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n# num_epochs = 80  # Number of epochs to train for\n# # scheduler=StepLR(optimizer, step_size=3, gamma=0.1)\n# validation_results,globalmodellll,Best_Clients_Model=train(client_models,globalmodel, clients_traindataloader, clients_valdataloader,client_optimizers, client_scheduler, num_epochs, device)\n\nnum_epochs = 45  # Number of epochs to train for\n\n# Run clustered federated training\nvalidation_results, cluster_models, Best_Clients_Model , cluster_labels = train(\n    client_models,\n    globalmodel,\n    clients_traindataloader,\n    clients_valdataloader,\n    client_optimizers,\n    client_scheduler,\n    num_epochs,\n    device,\n    k=3  # or however many clusters you want\n)\n","metadata":{"execution":{"iopub.status.busy":"2025-10-05T12:19:55.307336Z","iopub.execute_input":"2025-10-05T12:19:55.307603Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\n\ndef plot_validation_results_separate(validation_results, num_epochs):\n    # Plot validation accuracy of one global model on full combined dataset (if exists)\n    plt.figure(figsize=(12, 5))\n    all_data_acc = [acc for _, acc in validation_results[0].get(\"All_data\", [])]\n    \n    if all_data_acc:\n        plt.plot(range(1, len(all_data_acc) + 1), all_data_acc, label=\"Global on Full Data\", color='blue')\n        plt.xlabel('Epoch')\n        plt.ylabel('Validation Accuracy (%)')\n        plt.title('Validation Accuracy of Global Model on Full Dataset')\n        plt.grid(True)\n        plt.legend()\n        plt.tight_layout()\n        plt.show()\n    else:\n        print(\"No 'All_data' validation results to plot.\")\n\n    # Plot validation accuracy for each client's model after aggregation\n    plt.figure(figsize=(14, 6))\n    for client_idx, results in validation_results.items():\n        after_acc = [acc for _, acc in results[\"after\"]]\n        if after_acc:\n            plt.plot(range(1, len(after_acc) + 1), after_acc, label=f'Client {client_idx}')\n    \n    plt.xlabel('Epoch')\n    plt.ylabel('Validation Accuracy (%)')\n    plt.title('Validation Accuracy of Cluster Models on Individual Clients')\n    plt.legend(loc='upper right', bbox_to_anchor=(1.15, 1), ncol=2, fontsize='small')\n    plt.grid(True)\n    plt.tight_layout()\n    plt.show()\nplot_validation_results_separate(validation_results, num_epochs)\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from collections import defaultdict\n\ndef evaluate_cluster_models_on_clients(cluster_models, cluster_labels, client_models, clients_testdataloader, device):\n    print(\"\\n--- Final Evaluation: Cluster Models on Corresponding Client Test Sets ---\\n\")\n    results = defaultdict(list)\n\n    for client_idx, cluster_id in enumerate(cluster_labels):\n        # Load the correct cluster model weights into the client model\n        model = client_models[client_idx]\n        cluster_weights = cluster_models[cluster_id]\n        state_dict = model.state_dict()\n        for key in cluster_weights:\n            state_dict[key] = cluster_weights[key]\n        model.load_state_dict(state_dict)\n        model.to(device)\n\n        # Evaluate on that client's test data\n        test_loader = clients_testdataloader[client_idx]\n        loss, acc, _, _ = validate(model, test_loader, device)\n        print(f\"Client {client_idx} [Cluster {cluster_id}] - Test Loss: {loss:.4f}, Accuracy: {acc:.2f}%\")\n        results[cluster_id].append((client_idx, loss, acc))\n\n        torch.cuda.empty_cache()\n\n    return results\n    \ntest_results = evaluate_cluster_models_on_clients(cluster_models, cluster_labels, client_models, clients_testdataloader, device)\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n# def evaluate_federated(globalmodelll, clients_val_dataloader, device):\n#     \"\"\"\n#     Evaluates the global model on the combined validation datasets of all clients.\n    \n#     Args:\n#         client_models (list): List of client models to evaluate.\n#         clients_val_dataloader (list): List of DataLoader objects for each client's validation set.\n#         device (torch.device): The device (CPU/GPU) for evaluation.\n\n#     Returns:\n#         avg_loss (float): The average loss over the combined dataset.\n#         accuracy (float): The accuracy over the combined dataset.\n#     \"\"\"\n\n#     for idx, val_dataloader in enumerate(clients_val_dataloader):\n#                 val_loss, val_accuracy,all_probs, all_labels = validate(globalmodelll, val_dataloader, device)\n#                 print(f\"client {idx} test loss= {val_loss} test accuracy ={val_accuracy}\")\n    \n#     # Combine all validation datasets\n    \n#     combined_dataset = ConcatDataset([dataloader.dataset for dataloader in clients_val_dataloader])\n#     combined_dataloader = DataLoader(combined_dataset, batch_size=16, shuffle=False)  # Adjust batch_size if needed\n#     # val_loss, val_accuracy,all_probs, all_labels = validate(globalmodelll, combined_dataloader, device)\n#     # print(f\"Global {idx} test loss= {val_loss} test accuracy ={val_accuracy}\")\n\n    \n#     # Use the global model (assuming the first model represents the global model)\n#     globalmodel.eval()  # Set to evaluation mode\n\n#     total_loss = 0.0\n#     correct_predictions = 0\n#     total_predictions = 0\n\n#     criterion = torch.nn.CrossEntropyLoss()  # Define loss function\n\n#     with torch.no_grad():\n#             for images, labels in tqdm(combined_dataloader, desc=\"Evaluating on Combined Test Set\"):\n#                 images = images.to(device)\n#                 labels = labels.to(device)\n    \n#                 # Forward pass\n#                 outputs = globalmodelll(images)\n#                 logits = outputs.logits\n    \n#                 # Compute the loss\n#                 loss = criterion(logits, labels)\n    \n#                 # Accumulate loss\n#                 total_loss += loss.item()\n    \n#                 # Get predictions (for accuracy calculation)\n#                 _, predicted = torch.max(logits, 1)\n#                 correct_predictions += (predicted == labels).sum().item()\n#                 total_predictions += labels.size(0)\n#             print(total_predictions)\n\n#     print(len(combined_dataloader))\n#     # Calculate average loss and accuracy\n#     avg_loss = total_loss / len(combined_dataloader)\n#     accuracy = (correct_predictions / total_predictions) * 100  # Accuracy in percentage\n\n#     return avg_loss, accuracy\n\n# # Example usage\n# test_loss, test_accuracy = evaluate_federated(client_models[9], clients_testdataloader, device)\n# print(f\"Test Loss: {test_loss:.4f}, Test Accuracy: {test_accuracy:.2f}%\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"torch.save(client_models[20].state_dict(), '/kaggle/working/globalmodel.pth')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from IPython.display import FileLink\nFileLink(r'/kaggle/working/globalmodel.pth')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for i in range(len(client_models)):\n                # client_models[i].load_state_dict(Best_Clients_Model[i]['model'])\n                client_models[i].load_state_dict(Best_Clients_Model[i]['model'])\n                client_models[i].to(device)\n                test_loss, test_accuracy = evaluate_federated(client_models[i], clients_testdataloader, device)\n                print(f\"Test Loss: {test_loss:.4f}, Test Accuracy: {test_accuracy:.2f}%\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for idx, val_dataloader in enumerate(clients_testdataloader):\n        val_loss, val_accuracy,all_probs, all_labels = validate(client_models[idx], val_dataloader, device)\n        print(f\"Client {idx + 1} Validation Loss After: {val_loss:.4f}, Validation Accuracy After: {val_accuracy:.2f}%\")\n            ","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}