{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.11","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":96164,"databundleVersionId":11418275,"sourceType":"competition"}],"dockerImageVersionId":31041,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"\n### Generating Decision Trees with GFlowNets: A Scalable Approach\n\nThe provided research outlines a sophisticated method called **DT-GFN** for creating decision tree models (Mahfoud et al., 2025). Instead of using traditional greedy algorithms, it frames the construction of a decision tree as a sequential planning problem and uses a deep reinforcement learning agent, specifically a Generative Flow Network (GFlowNet), to solve it (Mahfoud et al., 2025). The key innovation is training a generative model that samples diverse, high-quality decision trees from the Bayesian posterior distribution, which allows for principled ensembling and improved generalization (Mahfoud et al., 2025).\n\n#### 1. The Core Idea: Decision Trees as Sequential Decisions\n\nThe fundamental concept of DT-GFN is to model the construction of a decision tree as a trajectory within a Markov Decision Process (MDP) (Mahfoud et al., 2025).\n\n* **States (s)**: Each state in the MDP represents a partially built tree. The initial state, $s_0$, is an empty tree with just a root node (Mahfoud et al., 2025).\n* **Actions (a)**: An action consists of adding a new decision rule (a feature and a split threshold) to an existing leaf node, thereby expanding the tree (Mahfoud et al., 2025). Actions that would result in an invalid split (e.g., where no data points fall into a child node) are masked (Mahfoud et al., 2025).\n* **Trajectories ($\\tau$)**: A complete trajectory is a sequence of states and actions, $\\tau = (s_0 \\rightarrow s_1 \\rightarrow \\dots \\rightarrow s_n)$, that results in a finished decision tree, $T$ (Mahfoud et al., 2025).\n\nA deep neural network is trained as a **policy**, $P_F(s'|s)$, which learns the probability of transitioning to the next tree state $s'$ given the current partial tree $s$ (Mahfoud et al., 2025).\n\n#### 2. The Objective: Sampling from the Bayesian Posterior\n\nThe GFlowNet is trained to generate trees where the probability of sampling a specific tree, $T$, is proportional to a reward function, $R(T)$ (Mahfoud et al., 2025). DT-GFN defines this reward based on the principles of Bayesian inference, aiming to sample from the posterior distribution $\\mathbb{P}[T | \\text{Data}]$ (Mahfoud et al., 2025).\n\nThis reward function elegantly balances two factors:\n1.  **Data Likelihood $\\mathbb{P}[\\text{Data} | T]$**: How well the tree structure explains the provided training data. This is calculated using the Bayesian marginal likelihood (Mahfoud et al., 2025).\n2.  **Structural Prior $\\mathbb{P}[T]$**: A penalty based on the complexity of the tree, implementing Occam's razor. This prior favors trees with shorter description lengths, which is measured by the number of decision nodes (Mahfoud et al., 2025).\n\nBy optimizing this objective, DT-GFN learns to generate a diverse set of trees that are both accurate and interpretable (Mahfoud et al., 2025).\n\n#### 3. The Challenge: Scaling to Deep Trees\n\nA primary challenge in training GFlowNets is the \"credit assignment\" problem, especially on tasks that require long action sequences, like building deep decision trees (Zhang et al., 2023). Standard training objectives like **Trajectory Balance (TB)**, used by DT-GFN, compute a loss based on the entire trajectory (Mahfoud et al., 2025). This approach suffers from two major drawbacks for deep trees:\n\n* **Exploding Memory Costs**: The need to store activations for every step in a long trajectory makes training deep trees computationally expensive, quickly exceeding GPU memory limits (Zhang et al., 2023).\n* **Delayed Reward Signals**: The learning signal from the final reward must propagate backward through the entire long sequence, making it difficult for the agent to attribute credit to early actions ( Zhang et al., 2023).\n\n#### 4. The Solution: Efficient Training with Transition-Based Learning\n\nTo overcome these scaling issues, the literature proposes two critical training innovations, which are particularly relevant for making DT-GFN practical for complex problems  . These techniques, validated in the context of combinatorial optimization, enable efficient training on long-horizon tasks (Zhang et al., 2023).\n\n* **Transition-Based Training**: Instead of computing loss over entire trajectories, this approach treats individual transitions ($s \\rightarrow s'$) as training samples (Zhang et al., 2023). A buffer of transitions is collected from many rollouts, and the model is trained on mini-batches of these transitions ( ; Zhang et al., 2023). This decouples the GPU memory cost from the trajectory length, making it feasible to train very deep trees  .\n\n* **Forward-Looking (FL) Loss**: This technique provides a denser, intermediate learning signal at every step of the trajectory (Zhang et al., 2023). It augments the standard loss with a potential function $\\tilde{\\mathcal{E}}(s)$ that estimates the \"value\" of an intermediate state (e.g., based on the number of splits so far). This gives the agent immediate feedback on its actions, dramatically speeding up credit assignment and accelerating convergence (Zhang et al., 2023).\n\nBy combining the principled Bayesian framework of DT-GFN with these advanced, transition-based training techniques, the method evolves from a \"memory-limited research toy\" into a scalable and powerful model capable of tackling large-scale datasets.\n\n---\n### **References**\n\n1.  **Mahfoud, M., Boukachab, G., Koziarski, M., Hernandez-Garcia, A., Bauer, S., Bengio, Y., & Malkin, N. (2025).** *Learning Decision Trees as Amortized Structure Inference*. arXiv:2503.06985v1.\n2.  **Zhang, D., Dai, H., Malkin, N., Courville, A., Bengio, Y., & Pan, L. (2023).** *Let the Flows Tell: Solving Graph Combinatorial Optimization Problems with GFlowNets*. 37th Conference on Neural Information Processing Systems (NeurIPS).","metadata":{}},{"cell_type":"code","source":"%%writefile dtgfn_boosted_production.py\n# dtgfn_boosted_production.py – Production-Ready, Final GFN-Boost Pipeline\n# ------------------------------------------------------------------------------------------------\n# This version incorporates a full suite of fixes for show-stopper bugs, performance,\n# and stability, based on a detailed code review. This is the definitive version.\n# -------------------------------------------------------------------------------------------------\nfrom __future__ import annotations\nimport math, random, argparse\nfrom dataclasses import dataclass\nfrom pathlib import Path\nfrom typing import List, Tuple, Dict, Any, Set\nfrom collections import deque\nimport numpy as np\nimport pandas as pd\nimport torch, torch.nn as nn\nfrom tqdm import tqdm\nimport lightgbm as lgb\nfrom torch.optim.lr_scheduler import LambdaLR, CosineAnnealingLR, SequentialLR\n\n# ---------------------- 0. CLI Args ----------------------------------\ndef parse_args():\n    \"\"\"Parses command-line arguments.\"\"\"\n    p = argparse.ArgumentParser(description=\"Train a production-ready Boosted DT-GFN model.\")\n    p.add_argument(\"--boosting-lr\", type=float, default=0.1, help=\"Learning rate for the boosting updates.\")\n    p.add_argument(\"--top-k-trees\", type=int, default=10, help=\"Number of best trees to average at each boosting step.\")\n    p.add_argument(\"train\", type=Path, help=\"Path to training data.\")\n    p.add_argument(\"test\",  type=Path, help=\"Path to test data.\")\n    p.add_argument(\"--out\",       type=Path, default=\"drw_preds.csv\", help=\"Output path for predictions.\")\n    p.add_argument(\"--device\",    default=\"cuda\", help=\"Device for training.\")\n    p.add_argument(\"--updates\",   type=int, default=50, help=\"Number of boosting rounds.\")\n    p.add_argument(\"--rollouts\",  type=int, default=60, help=\"On-policy trajectories per boosting round.\")\n    p.add_argument(\"--batch\",     type=int, default=8192, help=\"Batch size for GFN evaluation.\")\n    p.add_argument(\"--bins\",      type=int, default=255, help=\"Number of feature bins.\")\n    p.add_argument(\"--lstm-hidden\", type=int, default=512, help=\"Dimension of the LSTM state tracker.\")\n    p.add_argument(\"--mlp-layers\", type=int, default=12, help=\"Number of hidden layers in MLP heads.\")\n    p.add_argument(\"--mlp-width\", type=int, default=256, help=\"Width of hidden layers in MLP heads.\")\n    p.add_argument(\"--max-depth\", type=int, default=7, help=\"Maximum depth for generated trees.\")\n    p.add_argument(\"--lr\", type=float, default=5e-5, help=\"Peak learning rate for the GFN policy network schedule.\")\n    p.add_argument(\"--beta-start\", type=float, default=0.35, help=\"Starting value for beta (split penalty).\")\n    p.add_argument(\"--beta-end\", type=float, default=math.log(4), help=\"Final value for beta (split penalty).\")\n    p.add_argument(\"--prior-scale\", type=float, default=0.5, help=\"Scaling factor for the gain-based prior.\")\n    return p.parse_args()\n\n# ---------------------- 1. Vocab & Tok -------------------------------\n@dataclass\nclass Vocab:\n    num_feat: int; num_th: int; num_leaf: int\n    PAD: int=0; BOS: int=1; EOS: int=2\n    @property\n    def split_start(self) -> int: return 3\n    def size(self) -> int: return self.split_start + self.num_feat + self.num_th + self.num_leaf\n\nclass Tok:\n    def __init__(self, v: Vocab): self.v = v\n    def _feat(self, i: int) -> int: return self.v.split_start + i\n    def _th(self, i: int) -> int: return self.v.split_start + self.v.num_feat + i\n    def _leaf(self, i: int) -> int: return self.v.split_start + self.v.num_feat + self.v.num_th + i\n    def decode_one(self, tid: int) -> Tuple[str, int]:\n        if tid < self.v.split_start: raise ValueError(f\"Invalid token id: {tid}\")\n        rem = tid - self.v.split_start\n        if rem < self.v.num_feat: return \"feat\", rem\n        rem -= self.v.num_feat\n        if rem < self.v.num_th: return \"th\", rem\n        return \"leaf\", rem - self.v.num_th\n    def decode(self, ids: List[int]) -> List[Tuple[str, int]]:\n        return [self.decode_one(i) for i in ids if i not in (self.v.BOS, self.v.EOS)]\n\n# ---------------------- 2. Policy Architecture (LSTM) -------------------------\n# No decorator on the class itself\nclass PolicyPaperMLP(nn.Module):\n    def __init__(self, vocab_sz: int, lstm_hidden: int, mlp_layers: int, mlp_width: int):\n        super().__init__()\n        self.embedding = nn.Embedding(vocab_sz, lstm_hidden)\n        self.rnn = nn.LSTM(input_size=lstm_hidden, hidden_size=lstm_hidden, num_layers=1, batch_first=True)\n        layers = [nn.Linear(lstm_hidden, mlp_width), nn.ReLU()]\n        for _ in range(mlp_layers - 1):\n            layers.extend([nn.Linear(mlp_width, mlp_width), nn.ReLU()])\n        self.shared_mlp = nn.Sequential(*layers)\n        self.head_tok = nn.Linear(mlp_width, vocab_sz)\n        self.head_flow = nn.Linear(mlp_width, 1)\n\n    def forward(self, seq: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:\n        rnn_output, _ = self.rnn(self.embedding(seq))\n        shared_output = self.shared_mlp(rnn_output)\n        return self.head_tok(shared_output), self.head_flow(shared_output).squeeze(-1)\n\n    @torch.jit.export  # <--- ADD THIS DECORATOR\n    def log_prob(self, seq: torch.Tensor) -> torch.Tensor:\n        if seq.size(1) < 2: return torch.empty(seq.size(0), 0, device=seq.device)\n        logits, _ = self.forward(seq[:, :-1]); return torch.gather(logits.log_softmax(dim=-1), -1, seq[:, 1:].unsqueeze(-1)).squeeze(-1)\n\n    @torch.jit.export  # <--- ADD THIS DECORATOR\n    def log_F(self, seq: torch.Tensor) -> torch.Tensor:\n        _, flow_values = self.forward(seq); return flow_values\n\n# ---------------------- 3. Losses, Env, and Helpers -------------------\n@torch.jit.script\ndef tb_loss(log_pf, log_pb, log_z, R, prior):\n    return ((log_z + log_pf.sum(1) - (torch.log(R) + prior + log_pb.sum(1)))**2).mean()\n\n@torch.jit.script\ndef fl_loss(logF, log_pf, log_pb, dE):\n    return ((logF[:, :-1] + log_pf - (logF[:, 1:] + log_pb - dE))**2).mean()\n\nclass DRWEnv:\n    def __init__(self, df, feats, target_col, bins, device):\n        self.device = device\n        self.X_full = self._featurise(df, df, feats, bins)\n        self.y_full = torch.tensor(df[target_col].values, dtype=torch.float32, device=device)\n        self.y = self.y_full.clone()\n    def _featurise(self, df_target, df_source, feats, bins):\n        X_binned = []\n        for f in feats:\n            s = df_source[f].replace([np.inf,-np.inf],np.nan).fillna(df_source[f].median()).values\n            qs = np.linspace(0,1,bins+1); edges = np.quantile(s,qs)\n            edges = np.unique(edges); edges[0] -= 1e-9; edges[-1] += 1e-9\n            s_eval = df_target[f].replace([np.inf,-np.inf],np.nan).fillna(df_source[f].median()).values\n            X_binned.append(np.searchsorted(edges,s_eval,side=\"right\")-1)\n        return torch.tensor(np.stack(X_binned,1).astype(np.int32), device=self.device)\n    def reset(self, n: int):\n        self.idxs = torch.from_numpy(np.random.choice(len(self.y),n,replace=False)).to(self.device)\n        self.paths, self.open_leaves, self.done = [], 1, False\n    def step(self, action: Tuple[str, int]):\n        self.paths.append(action); kind, _ = action\n        if kind == \"feat\": self.open_leaves += 1\n        elif kind == \"leaf\": self.open_leaves -= 1\n        self.done = (self.open_leaves == 0 or len(self.paths) > 8192)\n    def evaluate(self, current_beta):\n        prior = -current_beta * sum(1 for k, _ in self.paths if k == \"feat\")\n        y_t = self.y[self.idxs]\n        if y_t.numel() == 0: return 0.0, torch.tensor([prior], device=self.device), None, None\n        X_batch = self.X_full[self.idxs]\n        base_mse = ((y_t - y_t.mean())**2).mean()\n        if not self.paths or self.open_leaves != 0:\n            return 1 / (1 + base_mse), torch.tensor([prior], device=self.device), None, y_t\n        path_iter = iter(self.paths)\n        def build():\n            try: k,i = next(path_iter)\n            except StopIteration: return None\n            if k == \"feat\": return {\"f\": i, \"t\": next(path_iter)[1], \"L\": build(), \"R\": build()}\n            return \"leaf_node\"\n        tree = build()\n        if not tree: return 1 / (1 + base_mse), torch.tensor([prior], device=self.device), None, y_t\n        pred = torch.empty_like(y_t)\n        stack = [(tree, torch.arange(y_t.numel(), device=self.device))]\n        while stack:\n            node, idx = stack.pop()\n            if not idx.numel() or node is None: continue\n            if isinstance(node, dict):\n                f,t,L,R = node['f'],node['t'],node['L'],node['R']\n                mask = X_batch[idx, f] <= t\n                stack.extend([(R, idx[~mask]), (L, idx[mask])])\n            else:\n                pred[idx] = y_t[idx].mean() if idx.numel() > 0 else y_t.mean()\n        mse = ((pred - y_t)**2).mean()\n        return 1 / (1 + mse), torch.tensor([prior], device=self.device), pred, y_t\n\ndef deltaE_split_gain(tokens, tok, env):\n    y = env.y[env.idxs]; N = y.numel()\n    dE = torch.zeros(tokens.shape[1] - 1, device=y.device)\n    if N > 0:\n        y2 = y * y; full_mse = (y2.mean() - y.mean() ** 2).item()\n    else: full_mse = 0.0\n    stack_rows = [torch.arange(N, device=y.device)]; stack_mse = [full_mse]\n    action_sequence = tokens[0, 1:-1].tolist(); it = iter(tok.decode(action_sequence))\n    token_idx = 0\n    for kind, idx in it:\n        if kind == \"feat\":\n            token_idx += 1\n            try: _, th = next(it)\n            except StopIteration: break\n            if not stack_rows: break\n            parent_rows = stack_rows.pop(); parent_mse = stack_mse.pop()\n            fv = env.X_full[env.idxs[parent_rows], idx]; mask = fv <= th\n            L_rows = parent_rows[mask]; R_rows = parent_rows[~mask]\n            def mse_fn(rows):\n                if rows.numel() < 2: return 0.0\n                yy = y[rows]; return ((yy*yy).mean() - yy.mean()**2).item()\n            mseL, mseR = mse_fn(L_rows), mse_fn(R_rows)\n            stack_rows.extend([R_rows, L_rows]); stack_mse.extend([mseR, mseL])\n            wL = L_rows.numel(); wR = R_rows.numel(); parent_N = wL + wR\n            if parent_N > 0:\n                gain = parent_mse - (wL/parent_N * mseL + wR/parent_N * mseR)\n                dE[token_idx] = -gain\n            token_idx += 1\n        else:\n            if stack_rows: stack_rows.pop(); stack_mse.pop()\n            token_idx += 1\n    return dE.unsqueeze(0)\n\ndef get_tree_predictor(traj, X_binned, y_target, tok):\n    path_iter = iter(tok.decode(traj[1:-1]))\n    def build_recursive():\n        try: k, i = next(path_iter)\n        except StopIteration: return None\n        if k == 'feat': return {'type': 'split', 'f': i, 't': next(path_iter)[1], 'L': build_recursive(), 'R': build_recursive()}\n        return {'type': 'leaf', 'value': 0}\n    tree_structure = build_recursive()\n    if tree_structure is None: return lambda X: torch.zeros(X.size(0), device=X.device)\n    q = [(tree_structure, torch.arange(X_binned.size(0), device=X_binned.device))]\n    while q:\n        node, indices = q.pop(0)\n        if node['type'] == 'split':\n            if not indices.numel() or node.get('L') is None: continue\n            mask = X_binned[indices, node['f']] <= node['t']\n            q.append((node['L'], indices[mask])); q.append((node['R'], indices[~mask]))\n        else:\n            node['value'] = y_target[indices].mean().item() if indices.numel() > 0 else y_target.mean().item()\n    def predict(X_test):\n        out = torch.empty(X_test.size(0), device=X_test.device)\n        stack = [(tree_structure, torch.arange(X_test.size(0), device=X_test.device))]\n        while stack:\n            node, indices = stack.pop()\n            if not indices.numel() or not node: continue\n            if node['type'] == 'leaf': out[indices] = node['value']\n            else:\n                mask = X_test[indices, node['f']] <= node['t']\n                if node.get('L'): stack.append((node['L'], indices[mask]))\n                if node.get('R'): stack.append((node['R'], indices[~mask]))\n        return out\n    return predict\n\nclass ReplayBuffer:\n    def __init__(self, capacity=10000):\n        self.capacity, self.data = capacity, []\n    def add(self, r, t, p, idxs):\n        self.data.append((r, t, p, idxs)); self.data.sort(key=lambda x:x[0], reverse=True)\n        if len(self.data) > self.capacity: self.data.pop()\n    def sample(self, k): return random.sample(self.data, min(k, len(self.data)))\n\ndef _safe_sample(logits, mask, temperature):\n    logits = logits / temperature; masked = torch.where(mask, logits, torch.tensor(-1e9, device=logits.device))\n    probs = torch.softmax(masked, dim=-1); return torch.multinomial(probs, 1).item()\n\ndef create_gain_bias(df_train, feats, target, tok, bins, prior_scale=0.5) -> torch.Tensor:\n    if len(df_train) > 200_000:\n        df_sample = df_train.sample(n=200_000, random_state=42)\n    else:\n        df_sample = df_train\n    X_binned = []\n    for f in feats:\n        s = df_sample[f].replace([np.inf,-np.inf],np.nan).fillna(df_sample[f].median()).values\n        qs = np.linspace(0,1,bins+1); edges = np.quantile(s,qs)\n        edges = np.unique(edges); edges[0] -= 1e-9; edges[-1] += 1e-9\n        X_binned.append(np.searchsorted(edges,s,side=\"right\")-1)\n    X_binned = np.stack(X_binned,1)\n    y_train = df_sample[target].values\n    lgb_train = lgb.Dataset(X_binned, y_train, feature_name=feats)\n    params = { 'objective': 'regression_l1', 'metric': 'l1', 'n_estimators': 100, 'learning_rate': 0.05, 'feature_fraction': 0.8, 'bagging_fraction': 0.8, 'bagging_freq': 1, 'num_leaves': 1024, 'max_depth': 10, 'verbose': -1, 'n_jobs': -1 }\n    gbm = lgb.train(params, lgb_train)\n    tree_info = gbm.dump_model()[\"tree_info\"]\n    all_splits = []\n    # This function is INSIDE create_gain_bias\n    def parse_node(node):\n        nonlocal bias  # <--- ADD THIS LINE\n        if \"split_gain\" in node and node[\"split_gain\"] > 0:\n            f_idx = node[\"split_feature\"]\n            gain = node[\"split_gain\"]\n            \n            tok_id_feat = tok._feat(f_idx)\n            if tok_id_feat < tok.v.size():\n                bias[tok_id_feat] += gain\n    \n            try:\n                bin_idx = min(int(node[\"threshold\"]), bins - 1)\n                tok_id_th = tok._th(bin_idx)\n                if tok_id_th < tok.v.size():\n                    bias[tok_id_th] += gain\n            except ValueError:\n                try:\n                    categories = [int(c) for c in node[\"threshold\"].split('||')]\n                    if len(categories) > 0:\n                        distributed_gain = gain / len(categories)\n                        for cat_idx in categories:\n                            bin_idx = min(cat_idx, bins - 1)\n                            tok_id_th = tok._th(bin_idx)\n                            if tok_id_th < tok.v.size():\n                                bias[tok_id_th] += distributed_gain\n                except (ValueError, AttributeError):\n                    pass\n    \n        if \"left_child\" in node: parse_node(node[\"left_child\"])\n        if \"right_child\" in node: parse_node(node[\"right_child\"])\n    gains = np.array([g for _, _, g in all_splits])\n    if len(gains) == 0: return torch.zeros(tok.v.size(), dtype=torch.float32)\n    mean, std = gains.mean(), gains.std()\n    is_valid = np.abs(gains - mean) < 3 * std\n    bias = np.zeros(tok.v.size())\n    for (f, bin_, gain), is_valid_gain in zip(all_splits, is_valid):\n        if is_valid_gain:\n            tok_id_feat, tok_id_th = tok._feat(f), tok._th(bin_)\n            if tok_id_feat < tok.v.size() and tok_id_th < tok.v.size():\n                bias[tok_id_feat] += gain; bias[tok_id_th] += gain\n    bias_std = bias.std()\n    if bias_std > 1e-6: bias /= bias_std\n    bias *= (prior_scale / 2.0)\n    return torch.tensor(bias, dtype=torch.float32)\n\n# ------------------------- 4. Main Training & Inference Logic ---------------------------\ndef train_and_ensemble(args: argparse.Namespace) -> None:\n    print(\"--- Setting up models and environment with PRODUCTION GFN-Boost Pipeline ---\")\n    df_tr, df_te = pd.read_parquet(args.train), pd.read_parquet(args.test)\n    \n    feats = sorted(list(set([\n        \"X863\", \"X856\", \"X598\", \"X862\", \"X385\", \"X852\", \"X603\", \"X860\", \"X674\",\n        \"X415\", \"X345\", \"X855\", \"X174\", \"X302\", \"X178\", \"X168\", \"X612\",\n        \"buy_qty\", \"sell_qty\", \"volume\", \"X888\", \"X421\", \"X333\", \"X292\",\n        \"bid_qty\", \"ask_qty\",\n        \"X344\", \"X137\", \"X532\"\n    ])))\n\n    v = Vocab(len(feats), args.bins, 1)\n    tok = Tok(v)\n    \n    env = DRWEnv(df_tr, feats, 'label', args.bins, args.device)\n    y_tr_true, X_tr_binned = env.y_full.clone(), env.X_full.clone()\n    \n    prior_bias = create_gain_bias(df_tr, feats, 'label', tok, args.bins, args.prior_scale).to(args.device)\n    \n    pf = PolicyPaperMLP(v.size(), args.lstm_hidden, args.mlp_layers, args.mlp_width).to(args.device)\n    pb = PolicyPaperMLP(v.size(), args.lstm_hidden, args.mlp_layers, args.mlp_width).to(args.device)\n    \n    pf = torch.jit.script(pf); pb = torch.jit.script(pb)\n    \n    with torch.no_grad():\n        pf.head_tok.bias.copy_(prior_bias)\n        pb.head_tok.bias.copy_(prior_bias)\n    \n    log_z = torch.zeros((), device=args.device, requires_grad=True)\n    optf, optb, optz = torch.optim.AdamW(pf.parameters(), lr=args.lr), torch.optim.AdamW(pb.parameters(), lr=args.lr), torch.optim.Adam([log_z], lr=args.lr/10)\n    optimizers = [optf, optb, optz]\n    warmup_updates, t_max_val = 10, max(1, args.updates - 10)\n    warmup_schedulers = [LambdaLR(opt, lr_lambda=lambda upd: min(1.0, upd / warmup_updates)) for opt in optimizers]\n    decay_schedulers = [CosineAnnealingLR(opt, T_max=t_max_val) for opt in optimizers]\n    schedulers = [SequentialLR(opt, schedulers=[ws, ds], milestones=[warmup_updates]) for opt, ws, ds in zip(optimizers, warmup_schedulers, decay_schedulers)]\n    buf = ReplayBuffer(capacity=10000)\n    \n    base_prediction = torch.full_like(y_tr_true, y_tr_true.mean())\n    boosting_ensemble = []\n    \n    BETA_ANNEAL_UPDATES, TEMP_ANNEAL_UPDATES, FL_LOSS_ANNEAL_UPDATES = 20.0, 20.0, 20.0\n    \n    print(\"--- Starting GFN-Boost Training ---\")\n    for upd in range(1, args.updates + 1):\n        residuals, env.y = y_tr_true - base_prediction, y_tr_true - base_prediction\n        progress = min(1.0, (upd - 1) / BETA_ANNEAL_UPDATES)\n        current_beta = 0.0#args.beta_start + progress * (args.beta_end - args.beta_start)\n        temperature = max(1.0, 5.0 - (upd - 1) * (4.0 / TEMP_ANNEAL_UPDATES))\n        lam_fl = 0.1#min(1.0, upd / FL_LOSS_ANNEAL_UPDATES)\n        tb_loss_acc, fl_loss_acc, complete_rollouts_this_update = 0, 0, 0\n        \n        for opt in optimizers: opt.zero_grad()\n        \n        pbar = tqdm(range(args.rollouts), desc=f\"Boosting Round {upd:02d}\", leave=False)\n        for _ in pbar:\n            env.reset(args.batch)\n            seq = [v.BOS]\n            open_leaf_depths = deque([0]) \n            with torch.no_grad():\n                while not env.done:\n                    if not open_leaf_depths: break\n                    x = torch.tensor([seq], device=args.device)\n                    logits, _ = pf.forward(x)\n                    logits = logits[0, -1]\n                    mask = torch.zeros_like(logits, dtype=torch.bool)\n                    current_depth = open_leaf_depths[-1]\n                    if env.open_leaves > 0 and current_depth < args.max_depth:\n                        mask[v.split_start : v.split_start + v.num_feat] = True\n                    if env.open_leaves > 0:\n                        mask[v.split_start + v.num_feat + v.num_th:] = True\n                    if not mask.any(): break\n                    tok1 = _safe_sample(logits, mask, temperature)\n                    kind, idx = tok.decode_one(tok1)\n                    if kind == 'feat':\n                        d = open_leaf_depths.pop(); open_leaf_depths.append(d + 1); open_leaf_depths.append(d + 1)\n                    else: open_leaf_depths.pop()\n                    env.step((kind, idx)); seq.append(tok1)\n                    if kind == 'feat':\n                        x2, _ = pf.forward(torch.tensor([seq], device=args.device))\n                        logits_th = x2[0, -1]\n                        th_mask = torch.zeros_like(logits_th, dtype=torch.bool)\n                        th_mask[v.split_start+v.num_feat : v.split_start+v.num_feat+v.num_th] = True\n                        tok2 = _safe_sample(logits_th, th_mask, temperature)\n                        env.step(tok.decode_one(tok2)); seq.append(tok2)\n            seq.append(v.EOS)\n\n            if env.open_leaves != 0: continue\n            complete_rollouts_this_update += 1\n            R, prior, _, _ = env.evaluate(current_beta)\n            buf.add(R.item(), seq, prior.item(), env.idxs.clone())\n            \n            tokens_fwd = torch.tensor([seq], device=args.device)\n            log_pf_fwd = pf.log_prob(tokens_fwd)\n            logF = pf.log_F(tokens_fwd)\n            tokens_bwd = torch.flip(tokens_fwd, dims=[1])\n            log_pb = pb.log_prob(tokens_bwd)\n            dE = deltaE_split_gain(tokens_fwd, tok, env)\n            \n            l_tb = tb_loss(log_pf_fwd, log_pb, log_z, R.unsqueeze(0), prior)\n            l_fl = fl_loss(logF, log_pf_fwd, log_pb, dE)\n            loss = l_tb + lam_fl * l_fl\n            loss.backward()\n            tb_loss_acc += l_tb.item(); fl_loss_acc += l_fl.item()\n\n        did_any_backward = complete_rollouts_this_update > 0\n        if did_any_backward:\n            denom = complete_rollouts_this_update\n            torch.nn.utils.clip_grad_norm_(pf.parameters(), 1.0)\n            torch.nn.utils.clip_grad_norm_(pb.parameters(), 1.0)\n            torch.nn.utils.clip_grad_value_([log_z], 1.0)\n            for opt in optimizers:\n                for group in opt.param_groups:\n                    for param in group['params']:\n                        if param.grad is not None: param.grad /= denom\n            for opt in optimizers: opt.step()\n        \n        if did_any_backward:\n            for scheduler in schedulers: scheduler.step()\n        \n        if buf.data:\n            top_k_trajectories = buf.data[:args.top_k_trees]\n            top_k_sequences = [t[1] for t in top_k_trajectories]\n            boosting_ensemble.append(top_k_sequences)\n            avg_residual_preds = torch.zeros_like(base_prediction)\n            for best_seq in top_k_sequences:\n                predictor = get_tree_predictor(best_seq, X_tr_binned, residuals, tok)\n                avg_residual_preds += predictor(X_tr_binned)\n            if len(top_k_sequences) > 0:\n                avg_residual_preds /= len(top_k_sequences)\n            \n            base_prediction += args.boosting_lr * avg_residual_preds\n            \n            final_corr = torch.corrcoef(torch.stack([base_prediction, y_tr_true]))[0,1].item()\n            print(f\"Boosting Round {upd:02d} | Comp: {complete_rollouts_this_update}/{args.rollouts} | \"\n                  f\"TB: {tb_loss_acc/denom} | FL: {fl_loss_acc/denom} | Train ρ: {final_corr:+.3f}\")\n        else:\n            print(f\"Boosting Round {upd:02d} | No complete trees generated.\")\n        \n        buf.data.clear()\n\n    print(\"\\n--- Training finished. Starting Boosted Inference. ---\")\n    X_te_binned = env._featurise(df_te, df_tr, feats, args.bins)\n    final_test_preds = torch.full((len(df_te),), y_tr_true.mean().item(), device=args.device)\n    inference_residuals = y_tr_true.clone()\n    for i, top_k_in_round in enumerate(tqdm(boosting_ensemble, desc=\"Building boosted ensemble\")):\n        avg_train_preds_for_round, avg_test_preds_for_round = torch.zeros_like(y_tr_true), torch.zeros_like(final_test_preds)\n        if not top_k_in_round: continue\n        \n        for tree_seq in top_k_in_round:\n            predictor = get_tree_predictor(tree_seq, X_tr_binned, inference_residuals, tok)\n            avg_train_preds_for_round += predictor(X_tr_binned)\n            avg_test_preds_for_round += predictor(X_te_binned)\n            \n        avg_train_preds_for_round /= len(top_k_in_round)\n        avg_test_preds_for_round /= len(top_k_in_round)\n\n        final_test_preds += args.boosting_lr * avg_test_preds_for_round\n        inference_residuals -= args.boosting_lr * avg_train_preds_for_round\n\n    try:\n        df_out = pd.read_csv('/kaggle/input/drw-crypto-market-prediction/sample_submission.csv')\n        df_out['prediction'] = final_test_preds.cpu().numpy()\n    except (FileNotFoundError, pd.errors.EmptyDataError):\n        df_out = pd.DataFrame({'id': df_te.get('id', range(len(df_te))), 'prediction': final_test_preds.cpu().numpy()})\n    df_out.to_csv(args.out, index=False)\n    print(f\"Inference complete. Predictions saved to -> {args.out}\")\n\nif __name__ == '__main__':\n    args = parse_args()\n    train_and_ensemble(args)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-01T03:05:43.260696Z","iopub.execute_input":"2025-07-01T03:05:43.261564Z","iopub.status.idle":"2025-07-01T03:05:43.278484Z","shell.execute_reply.started":"2025-07-01T03:05:43.261527Z","shell.execute_reply":"2025-07-01T03:05:43.277763Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# choose sane defaults for Kaggle's 16 GB GPU\n!python dtgfn_boosted_production.py /kaggle/input/drw-crypto-market-prediction/train.parquet /kaggle/input/drw-crypto-market-prediction/test.parquet \\\n    --device cuda    --updates 100   --rollouts 10   --batch 500_000  \n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-01T03:05:44.850680Z","iopub.execute_input":"2025-07-01T03:05:44.851343Z","iopub.status.idle":"2025-07-01T03:26:39.882191Z","shell.execute_reply.started":"2025-07-01T03:05:44.851311Z","shell.execute_reply":"2025-07-01T03:26:39.881448Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}