{
  "id": 556970,
  "title": "TabM with Pytorch Lightning",
  "url": "/competitions/jane-street-real-time-market-data-forecasting/discussion/556970",
  "author_name": "",
  "post_date": "2025-01-16T03:05:18.998347500Z",
  "votes": 3,
  "comment_count": 3,
  "views": 0,
  "content": "<p>TabM has shown good performance in some top solutions.<br>\nHere, I’m sharing the code for training and predicting TabM using the PyTorch Lightning framework:<br>\n<a href=\"https://www.kaggle.com/code/iwatatakuya/tabm-with-pytorch-lightning\" target=\"_blank\">https://www.kaggle.com/code/iwatatakuya/tabm-with-pytorch-lightning</a></p>\n<p>I hope this will be helpful for your future modeling endeavors!</p>",
  "messages": [
    {
      "id": "3098052",
      "postDate": "01/16/2025 03:05:19",
      "content": "<p>TabM has shown good performance in some top solutions.<br>\nHere, I’m sharing the code for training and predicting TabM using the PyTorch Lightning framework:<br>\n<a href=\"https://www.kaggle.com/code/iwatatakuya/tabm-with-pytorch-lightning\" target=\"_blank\">https://www.kaggle.com/code/iwatatakuya/tabm-with-pytorch-lightning</a></p>\n<p>I hope this will be helpful for your future modeling endeavors!</p>",
      "rawMarkdown": "TabM has shown good performance in some top solutions.\nHere, I’m sharing the code for training and predicting TabM using the PyTorch Lightning framework:\nhttps://www.kaggle.com/code/iwatatakuya/tabm-with-pytorch-lightning\n\nI hope this will be helpful for your future modeling endeavors!",
      "votes": null
    },
    {
      "id": "3098124",
      "postDate": "01/16/2025 05:37:38",
      "content": "<p>Hey, thanks for sharing! What score did you get on LB with TabM?</p>",
      "rawMarkdown": "Hey, thanks for sharing! What score did you get on LB with TabM?",
      "votes": null
    },
    {
      "id": "3098132",
      "postDate": "01/16/2025 05:57:44",
      "content": "<p>thanks for sharing!! big upvote</p>",
      "rawMarkdown": "thanks for sharing!! big upvote",
      "votes": null
    },
    {
      "id": "3098579",
      "postDate": "01/16/2025 16:46:30",
      "content": "<p>Thanks for the code but got negative values for val_r_square so tried code from<br>\ninference notebook and got .00512 after 4 epochs. May be caused by time             and symbol.</p>\n<p><a href=\"https://www.kaggle.com/code/i2nfinit3y/jane-street-tabm-ft-transformer-inference\" target=\"_blank\">https://www.kaggle.com/code/i2nfinit3y/jane-street-tabm-ft-transformer-inference</a></p>\n<p>Code for training is now at:<br>\n<a href=\"https://www.kaggle.com/code/i2nfinit3y/jane-street-tabm-training/notebook?scriptVersionId=217873214\" target=\"_blank\">https://www.kaggle.com/code/i2nfinit3y/jane-street-tabm-training/notebook?scriptVersionId=217873214</a></p>\n<p>from torch.utils.data import Dataset, DataLoader, TensorDataset</p>\n<p>delete w_y in =batch</p>\n<p>replace trainer.fit(model,… with:</p>\n<p>symbol_ids = df_tr.select('symbol_id').to_numpy()[:, 0]<br>\ntime_id = df_tr.select(\"time_id\").to_numpy()[0]<br>\ntimie_id_array = df_tr.select(\"time_id\").to_numpy()[:, 0]<br>\nX_test = df_tr[col_feature_list].to_numpy()<br>\nX_test_tensor = torch.tensor(X_test, dtype=torch.float32).to(device)<br>\nsymbol_tensor = torch.tensor(symbol_ids, dtype=torch.float32).to(device)<br>\ntime_tensor = torch.tensor(timie_id_array, dtype=torch.float32).to(device)<br>\nX_cat = X_test_tensor[:, [9, 10, 11]]<br>\nX_contt = X_test_tensor[:, [i for i in range(X_test_tensor.shape[1]) if i not in [9, 10, 11]]]<br>\nX_catt = (torch.concat([X_cat, symbol_tensor.unsqueeze(-1), time_tensor.unsqueeze(-1)], axis=1)).to(torch.int64)<br>\ny_train = df_tr.select(\"responder_6\" ).to_numpy().flatten()<br>\ny_traint = torch.tensor(y_train, dtype=torch.float32)<br>\nww = df_tr.select('weight').to_numpy().flatten()<br>\nwwt= torch.tensor(ww, dtype=torch.float32)</p>\n<p>symbol_ids = df_va.select('symbol_id').to_numpy()[:, 0]<br>\ntime_id = df_va.select(\"time_id\").to_numpy()[0]<br>\ntimie_id_array = df_va.select(\"time_id\").to_numpy()[:, 0]<br>\nX_test = df_va[col_feature_list].to_numpy()<br>\nX_test_tensor = torch.tensor(X_test, dtype=torch.float32).to(device)<br>\nsymbol_tensor = torch.tensor(symbol_ids, dtype=torch.float32).to(device)<br>\ntime_tensor = torch.tensor(timie_id_array, dtype=torch.float32).to(device)<br>\nX_cat = X_test_tensor[:, [9, 10, 11]]<br>\nX_contv = X_test_tensor[:, [i for i in range(X_test_tensor.shape[1]) if i not in [9, 10, 11]]]<br>\nX_catv = (torch.concat([X_cat, symbol_tensor.unsqueeze(-1), time_tensor.unsqueeze(-1)], axis=1)).to(torch.int64)<br>\ny_train = df_va.select(\"responder_6\" ).to_numpy().flatten()<br>\ny_trainv = torch.tensor(y_train, dtype=torch.float32)<br>\nww = df_va.select('weight').to_numpy().flatten()<br>\nwwv= torch.tensor(ww, dtype=torch.float32)</p>\n<p>dataset = TensorDataset(X_contt, X_catt,y_traint ,wwt)<br>\ndataloadert = DataLoader(dataset, batch_size=8192)</p>\n<p>dataset = TensorDataset(X_contv, X_catv,y_trainv ,wwv)<br>\ndataloaderv = DataLoader(dataset, batch_size=8192)</p>\n<p>trainer.fit(model, dataloadert, dataloaderv)</p>",
      "rawMarkdown": "Thanks for the code but got negative values for val_r_square so tried code from\ninference notebook and got .00512 after 4 epochs. May be caused by time             and symbol.\n\nhttps://www.kaggle.com/code/i2nfinit3y/jane-street-tabm-ft-transformer-inference\n\nCode for training is now at:\nhttps://www.kaggle.com/code/i2nfinit3y/jane-street-tabm-training/notebook?scriptVersionId=217873214\n\nfrom torch.utils.data import Dataset, DataLoader, TensorDataset\n\ndelete w_y in =batch\n\nreplace trainer.fit(model,... with:\n\nsymbol_ids = df_tr.select('symbol_id').to_numpy()[:, 0]\ntime_id = df_tr.select(\"time_id\").to_numpy()[0]\ntimie_id_array = df_tr.select(\"time_id\").to_numpy()[:, 0]\nX_test = df_tr[col_feature_list].to_numpy()\nX_test_tensor = torch.tensor(X_test, dtype=torch.float32).to(device)\nsymbol_tensor = torch.tensor(symbol_ids, dtype=torch.float32).to(device)\ntime_tensor = torch.tensor(timie_id_array, dtype=torch.float32).to(device)\nX_cat = X_test_tensor[:, [9, 10, 11]]\nX_contt = X_test_tensor[:, [i for i in range(X_test_tensor.shape[1]) if i not in [9, 10, 11]]]\nX_catt = (torch.concat([X_cat, symbol_tensor.unsqueeze(-1), time_tensor.unsqueeze(-1)], axis=1)).to(torch.int64)\ny_train = df_tr.select(\"responder_6\" ).to_numpy().flatten()\ny_traint = torch.tensor(y_train, dtype=torch.float32)\nww = df_tr.select('weight').to_numpy().flatten()\nwwt= torch.tensor(ww, dtype=torch.float32)\n\nsymbol_ids = df_va.select('symbol_id').to_numpy()[:, 0]\ntime_id = df_va.select(\"time_id\").to_numpy()[0]\ntimie_id_array = df_va.select(\"time_id\").to_numpy()[:, 0]\nX_test = df_va[col_feature_list].to_numpy()\nX_test_tensor = torch.tensor(X_test, dtype=torch.float32).to(device)\nsymbol_tensor = torch.tensor(symbol_ids, dtype=torch.float32).to(device)\ntime_tensor = torch.tensor(timie_id_array, dtype=torch.float32).to(device)\nX_cat = X_test_tensor[:, [9, 10, 11]]\nX_contv = X_test_tensor[:, [i for i in range(X_test_tensor.shape[1]) if i not in [9, 10, 11]]]\nX_catv = (torch.concat([X_cat, symbol_tensor.unsqueeze(-1), time_tensor.unsqueeze(-1)], axis=1)).to(torch.int64)\ny_train = df_va.select(\"responder_6\" ).to_numpy().flatten()\ny_trainv = torch.tensor(y_train, dtype=torch.float32)\nww = df_va.select('weight').to_numpy().flatten()\nwwv= torch.tensor(ww, dtype=torch.float32)\n\ndataset = TensorDataset(X_contt, X_catt,y_traint ,wwt)\ndataloadert = DataLoader(dataset, batch_size=8192)\n\ndataset = TensorDataset(X_contv, X_catv,y_trainv ,wwv)\ndataloaderv = DataLoader(dataset, batch_size=8192)\n\ntrainer.fit(model, dataloadert, dataloaderv)",
      "votes": null
    }
  ],
  "comments": [
    {
      "id": 3098124,
      "author_name": "alexeigor",
      "author_url": "",
      "post_date": "01/16/2025 05:37:38",
      "content": "<p>Hey, thanks for sharing! What score did you get on LB with TabM?</p>",
      "votes": null,
      "replies": []
    },
    {
      "id": 3098132,
      "author_name": "zoutain",
      "author_url": "",
      "post_date": "01/16/2025 05:57:44",
      "content": "<p>thanks for sharing!! big upvote</p>",
      "votes": null,
      "replies": []
    },
    {
      "id": 3098579,
      "author_name": "tomkkk",
      "author_url": "",
      "post_date": "01/16/2025 16:46:30",
      "content": "<p>Thanks for the code but got negative values for val_r_square so tried code from<br>\ninference notebook and got .00512 after 4 epochs. May be caused by time             and symbol.</p>\n<p><a href=\"https://www.kaggle.com/code/i2nfinit3y/jane-street-tabm-ft-transformer-inference\" target=\"_blank\">https://www.kaggle.com/code/i2nfinit3y/jane-street-tabm-ft-transformer-inference</a></p>\n<p>Code for training is now at:<br>\n<a href=\"https://www.kaggle.com/code/i2nfinit3y/jane-street-tabm-training/notebook?scriptVersionId=217873214\" target=\"_blank\">https://www.kaggle.com/code/i2nfinit3y/jane-street-tabm-training/notebook?scriptVersionId=217873214</a></p>\n<p>from torch.utils.data import Dataset, DataLoader, TensorDataset</p>\n<p>delete w_y in =batch</p>\n<p>replace trainer.fit(model,… with:</p>\n<p>symbol_ids = df_tr.select('symbol_id').to_numpy()[:, 0]<br>\ntime_id = df_tr.select(\"time_id\").to_numpy()[0]<br>\ntimie_id_array = df_tr.select(\"time_id\").to_numpy()[:, 0]<br>\nX_test = df_tr[col_feature_list].to_numpy()<br>\nX_test_tensor = torch.tensor(X_test, dtype=torch.float32).to(device)<br>\nsymbol_tensor = torch.tensor(symbol_ids, dtype=torch.float32).to(device)<br>\ntime_tensor = torch.tensor(timie_id_array, dtype=torch.float32).to(device)<br>\nX_cat = X_test_tensor[:, [9, 10, 11]]<br>\nX_contt = X_test_tensor[:, [i for i in range(X_test_tensor.shape[1]) if i not in [9, 10, 11]]]<br>\nX_catt = (torch.concat([X_cat, symbol_tensor.unsqueeze(-1), time_tensor.unsqueeze(-1)], axis=1)).to(torch.int64)<br>\ny_train = df_tr.select(\"responder_6\" ).to_numpy().flatten()<br>\ny_traint = torch.tensor(y_train, dtype=torch.float32)<br>\nww = df_tr.select('weight').to_numpy().flatten()<br>\nwwt= torch.tensor(ww, dtype=torch.float32)</p>\n<p>symbol_ids = df_va.select('symbol_id').to_numpy()[:, 0]<br>\ntime_id = df_va.select(\"time_id\").to_numpy()[0]<br>\ntimie_id_array = df_va.select(\"time_id\").to_numpy()[:, 0]<br>\nX_test = df_va[col_feature_list].to_numpy()<br>\nX_test_tensor = torch.tensor(X_test, dtype=torch.float32).to(device)<br>\nsymbol_tensor = torch.tensor(symbol_ids, dtype=torch.float32).to(device)<br>\ntime_tensor = torch.tensor(timie_id_array, dtype=torch.float32).to(device)<br>\nX_cat = X_test_tensor[:, [9, 10, 11]]<br>\nX_contv = X_test_tensor[:, [i for i in range(X_test_tensor.shape[1]) if i not in [9, 10, 11]]]<br>\nX_catv = (torch.concat([X_cat, symbol_tensor.unsqueeze(-1), time_tensor.unsqueeze(-1)], axis=1)).to(torch.int64)<br>\ny_train = df_va.select(\"responder_6\" ).to_numpy().flatten()<br>\ny_trainv = torch.tensor(y_train, dtype=torch.float32)<br>\nww = df_va.select('weight').to_numpy().flatten()<br>\nwwv= torch.tensor(ww, dtype=torch.float32)</p>\n<p>dataset = TensorDataset(X_contt, X_catt,y_traint ,wwt)<br>\ndataloadert = DataLoader(dataset, batch_size=8192)</p>\n<p>dataset = TensorDataset(X_contv, X_catv,y_trainv ,wwv)<br>\ndataloaderv = DataLoader(dataset, batch_size=8192)</p>\n<p>trainer.fit(model, dataloadert, dataloaderv)</p>",
      "votes": null,
      "replies": []
    }
  ],
  "raw_markdown_by_id": {
    "3098052": "TabM has shown good performance in some top solutions.\nHere, I’m sharing the code for training and predicting TabM using the PyTorch Lightning framework:\nhttps://www.kaggle.com/code/iwatatakuya/tabm-with-pytorch-lightning\n\nI hope this will be helpful for your future modeling endeavors!",
    "3098124": "Hey, thanks for sharing! What score did you get on LB with TabM?",
    "3098132": "thanks for sharing!! big upvote",
    "3098579": "Thanks for the code but got negative values for val_r_square so tried code from\ninference notebook and got .00512 after 4 epochs. May be caused by time             and symbol.\n\nhttps://www.kaggle.com/code/i2nfinit3y/jane-street-tabm-ft-transformer-inference\n\nCode for training is now at:\nhttps://www.kaggle.com/code/i2nfinit3y/jane-street-tabm-training/notebook?scriptVersionId=217873214\n\nfrom torch.utils.data import Dataset, DataLoader, TensorDataset\n\ndelete w_y in =batch\n\nreplace trainer.fit(model,... with:\n\nsymbol_ids = df_tr.select('symbol_id').to_numpy()[:, 0]\ntime_id = df_tr.select(\"time_id\").to_numpy()[0]\ntimie_id_array = df_tr.select(\"time_id\").to_numpy()[:, 0]\nX_test = df_tr[col_feature_list].to_numpy()\nX_test_tensor = torch.tensor(X_test, dtype=torch.float32).to(device)\nsymbol_tensor = torch.tensor(symbol_ids, dtype=torch.float32).to(device)\ntime_tensor = torch.tensor(timie_id_array, dtype=torch.float32).to(device)\nX_cat = X_test_tensor[:, [9, 10, 11]]\nX_contt = X_test_tensor[:, [i for i in range(X_test_tensor.shape[1]) if i not in [9, 10, 11]]]\nX_catt = (torch.concat([X_cat, symbol_tensor.unsqueeze(-1), time_tensor.unsqueeze(-1)], axis=1)).to(torch.int64)\ny_train = df_tr.select(\"responder_6\" ).to_numpy().flatten()\ny_traint = torch.tensor(y_train, dtype=torch.float32)\nww = df_tr.select('weight').to_numpy().flatten()\nwwt= torch.tensor(ww, dtype=torch.float32)\n\nsymbol_ids = df_va.select('symbol_id').to_numpy()[:, 0]\ntime_id = df_va.select(\"time_id\").to_numpy()[0]\ntimie_id_array = df_va.select(\"time_id\").to_numpy()[:, 0]\nX_test = df_va[col_feature_list].to_numpy()\nX_test_tensor = torch.tensor(X_test, dtype=torch.float32).to(device)\nsymbol_tensor = torch.tensor(symbol_ids, dtype=torch.float32).to(device)\ntime_tensor = torch.tensor(timie_id_array, dtype=torch.float32).to(device)\nX_cat = X_test_tensor[:, [9, 10, 11]]\nX_contv = X_test_tensor[:, [i for i in range(X_test_tensor.shape[1]) if i not in [9, 10, 11]]]\nX_catv = (torch.concat([X_cat, symbol_tensor.unsqueeze(-1), time_tensor.unsqueeze(-1)], axis=1)).to(torch.int64)\ny_train = df_va.select(\"responder_6\" ).to_numpy().flatten()\ny_trainv = torch.tensor(y_train, dtype=torch.float32)\nww = df_va.select('weight').to_numpy().flatten()\nwwv= torch.tensor(ww, dtype=torch.float32)\n\ndataset = TensorDataset(X_contt, X_catt,y_traint ,wwt)\ndataloadert = DataLoader(dataset, batch_size=8192)\n\ndataset = TensorDataset(X_contv, X_catv,y_trainv ,wwv)\ndataloaderv = DataLoader(dataset, batch_size=8192)\n\ntrainer.fit(model, dataloadert, dataloaderv)"
  },
  "source": "meta"
}