{
  "id": 575149,
  "title": "How to average the CV folds during training?",
  "url": "/competitions/byu-locating-bacterial-flagellar-motors-2025/discussion/575149",
  "author_name": "Rabia Mushtaq",
  "post_date": "2025-04-26T10:47:17.823000",
  "votes": 1,
  "comment_count": 4,
  "views": 0,
  "content": "<p>Newbie here! This might be a dumb question but <br>\nCan someone tell me how to average the Cv i am using 5 fold Cv during training but i am stuck at how do i ensemble or may i say average them? Do we need to average the model weights? Or should i ensemble during inference? Or is there something else? Would be a great help if someone could tell me what is the best option and how to do it properly</p>",
  "messages": [
    {
      "id": 3187650,
      "postDate": "2025-04-26T10:47:17.823Z",
      "content": "<p>Newbie here! This might be a dumb question but <br>\nCan someone tell me how to average the Cv i am using 5 fold Cv during training but i am stuck at how do i ensemble or may i say average them? Do we need to average the model weights? Or should i ensemble during inference? Or is there something else? Would be a great help if someone could tell me what is the best option and how to do it properly</p>",
      "rawMarkdown": "Newbie here! This might be a dumb question but \nCan someone tell me how to average the Cv i am using 5 fold Cv during training but i am stuck at how do i ensemble or may i say average them? Do we need to average the model weights? Or should i ensemble during inference? Or is there something else? Would be a great help if someone could tell me what is the best option and how to do it properly",
      "votes": 1
    },
    {
      "id": 3187723,
      "postDate": "2025-04-26T12:39:29.647Z",
      "content": "<p>like this </p>\n<pre><code>def (model_class, checkpoint_path, device):\n    model = ()\n    model.(torch.(checkpoint_path, map_location=device))\n    model.(device)\n    model.()\n    return model\n\ndef (model_class, checkpoint_paths, dataloader, device=):\n    models = [(model_class, ckpt, device) for ckpt in checkpoint_paths]\n\n    all_preds = []\n\n    with torch.():\n        for batch in (dataloader, desc=):\n            inputs = batch.(device)\n\n            fold_preds = []\n            for model in models:\n                preds = (inputs)\n                fold_preds.(preds)\n\n            # Stack along new dimension and average\n            fold_preds = torch.(fold_preds, dim=)  # (num_folds, batch_size, ...)\n            avg_preds = fold_preds.(dim=)  # (batch_size, ...)\n\n            all_preds.(avg_preds.())\n\n    all_preds = torch.(all_preds, dim=)\n    return all_preds\n</code></pre>",
      "rawMarkdown": "like this \n```\ndef load_model(model_class, checkpoint_path, device):\n    model = model_class()\n    model.load_state_dict(torch.load(checkpoint_path, map_location=device))\n    model.to(device)\n    model.eval()\n    return model\n\ndef fold_ensemble_predict(model_class, checkpoint_paths, dataloader, device='cuda'):\n    models = [load_model(model_class, ckpt, device) for ckpt in checkpoint_paths]\n    \n    all_preds = []\n    \n    with torch.no_grad():\n        for batch in tqdm(dataloader, desc=\"Ensemble Inference\"):\n            inputs = batch.to(device)\n\n            fold_preds = []\n            for model in models:\n                preds = model(inputs)\n                fold_preds.append(preds)\n\n            # Stack along new dimension and average\n            fold_preds = torch.stack(fold_preds, dim=0)  # (num_folds, batch_size, ...)\n            avg_preds = fold_preds.mean(dim=0)  # (batch_size, ...)\n            \n            all_preds.append(avg_preds.cpu())\n    \n    all_preds = torch.cat(all_preds, dim=0)\n    return all_preds\n```",
      "replies": [
        {
          "id": 3187731,
          "postDate": "2025-04-26T12:47:44.813Z",
          "content": "<p><a href=\"https://www.kaggle.com/seeingtimes\" target=\"_blank\">@seeingtimes</a> Thanks a lot! I suppose this snippet should be put in the inference notebook but will it not increase the inference time?</p>",
          "rawMarkdown": "@seeingtimes Thanks a lot! I suppose this snippet should be put in the inference notebook but will it not increase the inference time?",
          "replies": [
            {
              "id": 3188150,
              "postDate": "2025-04-27T05:41:30.740Z",
              "content": "<p>It will increase the inference time since you are inferring five models instead of one.</p>",
              "rawMarkdown": "It will increase the inference time since you are inferring five models instead of one."
            },
            {
              "id": 3188185,
              "postDate": "2025-04-27T07:23:35.703Z",
              "rawMarkdown": "",
              "isDeleted": true
            }
          ]
        }
      ]
    }
  ],
  "comments": [
    {
      "id": 3187723,
      "author_name": "Seeing Times",
      "author_url": "",
      "post_date": "2025-04-26T12:39:29.647000",
      "content": "<p>like this </p>\n<pre><code>def (model_class, checkpoint_path, device):\n    model = ()\n    model.(torch.(checkpoint_path, map_location=device))\n    model.(device)\n    model.()\n    return model\n\ndef (model_class, checkpoint_paths, dataloader, device=):\n    models = [(model_class, ckpt, device) for ckpt in checkpoint_paths]\n\n    all_preds = []\n\n    with torch.():\n        for batch in (dataloader, desc=):\n            inputs = batch.(device)\n\n            fold_preds = []\n            for model in models:\n                preds = (inputs)\n                fold_preds.(preds)\n\n            # Stack along new dimension and average\n            fold_preds = torch.(fold_preds, dim=)  # (num_folds, batch_size, ...)\n            avg_preds = fold_preds.(dim=)  # (batch_size, ...)\n\n            all_preds.(avg_preds.())\n\n    all_preds = torch.(all_preds, dim=)\n    return all_preds\n</code></pre>",
      "votes": 0,
      "replies": [
        {
          "id": 3187731,
          "author_name": "Rabia Mushtaq",
          "author_url": "",
          "post_date": "2025-04-26T12:47:44.813000",
          "content": "<p><a href=\"https://www.kaggle.com/seeingtimes\" target=\"_blank\">@seeingtimes</a> Thanks a lot! I suppose this snippet should be put in the inference notebook but will it not increase the inference time?</p>",
          "votes": 0,
          "replies": [
            {
              "id": 3188150,
              "author_name": "Seeing Times",
              "author_url": "",
              "post_date": "2025-04-27T05:41:30.740000",
              "content": "<p>It will increase the inference time since you are inferring five models instead of one.</p>",
              "votes": 0,
              "replies": []
            },
            {
              "id": 3188185,
              "author_name": "",
              "author_url": "",
              "post_date": "2025-04-27T07:23:35.703000",
              "content": "",
              "votes": 0,
              "replies": []
            }
          ]
        }
      ]
    }
  ],
  "raw_markdown_by_id": {
    "3187650": "Newbie here! This might be a dumb question but \nCan someone tell me how to average the Cv i am using 5 fold Cv during training but i am stuck at how do i ensemble or may i say average them? Do we need to average the model weights? Or should i ensemble during inference? Or is there something else? Would be a great help if someone could tell me what is the best option and how to do it properly",
    "3187723": "like this \n```\ndef load_model(model_class, checkpoint_path, device):\n    model = model_class()\n    model.load_state_dict(torch.load(checkpoint_path, map_location=device))\n    model.to(device)\n    model.eval()\n    return model\n\ndef fold_ensemble_predict(model_class, checkpoint_paths, dataloader, device='cuda'):\n    models = [load_model(model_class, ckpt, device) for ckpt in checkpoint_paths]\n    \n    all_preds = []\n    \n    with torch.no_grad():\n        for batch in tqdm(dataloader, desc=\"Ensemble Inference\"):\n            inputs = batch.to(device)\n\n            fold_preds = []\n            for model in models:\n                preds = model(inputs)\n                fold_preds.append(preds)\n\n            # Stack along new dimension and average\n            fold_preds = torch.stack(fold_preds, dim=0)  # (num_folds, batch_size, ...)\n            avg_preds = fold_preds.mean(dim=0)  # (batch_size, ...)\n            \n            all_preds.append(avg_preds.cpu())\n    \n    all_preds = torch.cat(all_preds, dim=0)\n    return all_preds\n```"
  }
}