-
Notifications
You must be signed in to change notification settings - Fork 7
Add regression-only Nori-Rel 30M #17
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
Closed
minkyu-choi07
wants to merge
3
commits into
PriorLabs:main
from
minkyu-choi07:feat/nori-rel-30m-regression
Closed
Changes from all commits
Commits
File filter
Filter by extension
Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
There are no files selected for viewing
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -0,0 +1,308 @@ | ||
| { | ||
| "cells": [ | ||
| { | ||
| "cell_type": "markdown", | ||
| "id": "7fb27b941602401d91542211134fc71a", | ||
| "metadata": {}, | ||
| "source": [ | ||
| "# Run Nori-Rel on one RelBench task\n", | ||
| "\n", | ||
| "This notebook walks through the complete Nori-Rel workflow: verify the environment, choose a supported regression task, precompute its relational features, run the frozen public Nori 30M checkpoint, and save the result.\n", | ||
| "\n", | ||
| "Nori-Rel performs in-context regression over depth-2 Deep Feature Synthesis (DFS) features. It does **not** fine-tune the checkpoint. The adapter is regression-only and deliberately disables silent context subsampling and cache quantization.\n", | ||
| "\n", | ||
| "> **Resources:** an NVIDIA GPU is strongly recommended. DFS runs on CPU and can be slow; larger datasets can also need substantial host RAM while the BF16 context cache is active. The first run downloads the selected RelBench database and the public checkpoint." | ||
| ] | ||
| }, | ||
| { | ||
| "cell_type": "markdown", | ||
| "id": "acae54e37e7d407bbb7b55eff062a284", | ||
| "metadata": {}, | ||
| "source": [ | ||
| "## Step 1 — install and launch the notebook\n", | ||
| "\n", | ||
| "Use Python 3.11 or 3.12. From a RelArena source checkout, launch Jupyter in the project environment:\n", | ||
| "\n", | ||
| "```bash\n", | ||
| "uv sync --all-packages --no-group kurversc --extra nori-rel\n", | ||
| "uv run --all-packages --no-sync --with jupyter jupyter lab examples/nori_rel.ipynb\n", | ||
| "```\n", | ||
| "\n", | ||
| "If Jupyter is already running from that environment, continue with the next cell." | ||
| ] | ||
| }, | ||
| { | ||
| "cell_type": "code", | ||
| "execution_count": null, | ||
| "id": "9a63283cbaf04dbcab1f6479b197f3a8", | ||
| "metadata": {}, | ||
| "outputs": [], | ||
| "source": [ | ||
| "import sys\n", | ||
| "from importlib.util import find_spec\n", | ||
| "\n", | ||
| "import torch\n", | ||
| "\n", | ||
| "assert sys.version_info[:2] in {(3, 11), (3, 12)}, sys.version\n", | ||
| "assert find_spec(\"synthefy_nori\") is not None, (\n", | ||
| " \"The nori-rel extra is missing. Relaunch with: \"\n", | ||
| " \"uv run --all-packages --no-sync --with jupyter jupyter lab\"\n", | ||
| ")\n", | ||
| "\n", | ||
| "device = torch.cuda.get_device_name(0) if torch.cuda.is_available() else \"CPU\"\n", | ||
| "print(f\"Python {sys.version.split()[0]} | torch {torch.__version__} | {device}\")\n", | ||
| "if not torch.cuda.is_available():\n", | ||
| " print(\"Warning: Nori-Rel can run on CPU, but a GPU is strongly recommended.\")" | ||
| ] | ||
| }, | ||
| { | ||
| "cell_type": "markdown", | ||
| "id": "8dd0d8092fe74a7c96281538738b07e2", | ||
| "metadata": {}, | ||
| "source": [ | ||
| "## Step 2 — choose storage and a task\n", | ||
| "\n", | ||
| "The example uses `rel-f1/driver-position`, a regression task. Change `DATASET` and `TASK` to another supported pair listed in the next step. Keep `SEED = 0` and `N_TRIALS = 1` to reproduce the submitted fixed configuration.\n", | ||
| "\n", | ||
| "Feature artifacts, downloaded data, and the checkpoint are kept outside the repository. Set `WORK_DIR` to a fast disk with enough free space." | ||
| ] | ||
| }, | ||
| { | ||
| "cell_type": "code", | ||
| "execution_count": null, | ||
| "id": "72eea5119410473aa328ad9291626812", | ||
| "metadata": {}, | ||
| "outputs": [], | ||
| "source": [ | ||
| "import os\n", | ||
| "from pathlib import Path\n", | ||
| "\n", | ||
| "DATASET = \"rel-f1\"\n", | ||
| "TASK = \"driver-position\"\n", | ||
| "SEED = 0\n", | ||
| "N_TRIALS = 1\n", | ||
| "\n", | ||
| "WORK_DIR = Path.home() / \".cache\" / \"relarena\" / \"nori-rel\"\n", | ||
| "FEATURE_CACHE = WORK_DIR / \"features\"\n", | ||
| "RESULTS_DIR = Path(\"results\") / \"nori-rel\"\n", | ||
| "\n", | ||
| "os.environ.setdefault(\"HF_HOME\", str(WORK_DIR / \"huggingface\"))\n", | ||
| "os.environ.setdefault(\"RELBENCH_CACHE_DIR\", str(WORK_DIR / \"relbench\"))\n", | ||
| "FEATURE_CACHE.mkdir(parents=True, exist_ok=True)\n", | ||
| "RESULTS_DIR.mkdir(parents=True, exist_ok=True)\n", | ||
| "\n", | ||
| "print(f\"Feature cache: {FEATURE_CACHE}\")\n", | ||
| "print(f\"Results: {RESULTS_DIR.resolve()}\")" | ||
| ] | ||
| }, | ||
| { | ||
| "cell_type": "markdown", | ||
| "id": "8edb47106e1a46a883d545849b8ab81b", | ||
| "metadata": {}, | ||
| "source": [ | ||
| "## Step 3 — confirm the task is supported\n", | ||
| "\n", | ||
| "Nori-Rel supports regression only. Listing the eligible RelBench tasks reads registry metadata and does not download any datasets." | ||
| ] | ||
| }, | ||
| { | ||
| "cell_type": "code", | ||
| "execution_count": null, | ||
| "id": "10185d26023b46108eb7d9f57d49d2b3", | ||
| "metadata": {}, | ||
| "outputs": [], | ||
| "source": [ | ||
| "import pandas as pd\n", | ||
| "from relbench.base import TaskType\n", | ||
| "\n", | ||
| "from relarena import list_entity_tasks\n", | ||
| "\n", | ||
| "regression_tasks = list_entity_tasks(task_types=frozenset({TaskType.REGRESSION}))\n", | ||
| "task_table = pd.DataFrame(\n", | ||
| " [(spec.dataset, spec.task) for spec in regression_tasks],\n", | ||
| " columns=[\"dataset\", \"task\"],\n", | ||
| ")\n", | ||
| "assert (DATASET, TASK) in {(spec.dataset, spec.task) for spec in regression_tasks}, (\n", | ||
| " f\"{DATASET}/{TASK} is not a supported regression task\"\n", | ||
| ")\n", | ||
| "task_table" | ||
| ] | ||
| }, | ||
| { | ||
| "cell_type": "markdown", | ||
| "id": "8763a12b2bbd4a93a75aff182afb95dc", | ||
| "metadata": {}, | ||
| "source": [ | ||
| "## Step 4 — warm the shared DFS feature cache\n", | ||
| "\n", | ||
| "Nori-Rel has no cache of its own. It reads the same leak-safe depth-2 DFS\n", | ||
| "matrices as RDBLearn and TabPFN-Rel, under the same cache keys, so one warmed\n", | ||
| "directory serves all three and `relarena.featurization.warm_cache` is the warmer\n", | ||
| "for all of them. If you already have a compatible cache, point `FEATURE_CACHE` at\n", | ||
| "it and skip this cell.\n", | ||
| "\n", | ||
| "The warmer fills the inner validation split and both outer-fit histories. This is\n", | ||
| "the CPU-heavy stage; run it once per dataset/task/cache combination and later\n", | ||
| "model runs read the stored Parquet artifacts. A cache miss during the experiment\n", | ||
| "is an error by design.\n" | ||
| ] | ||
| }, | ||
| { | ||
| "cell_type": "code", | ||
| "execution_count": null, | ||
| "id": "7623eae2785240b9bd12b16a66d81610", | ||
| "metadata": {}, | ||
| "outputs": [], | ||
| "source": [ | ||
| "from relarena import CacheConfig, RelBenchDatasetTask\n", | ||
| "from relarena.featurization.warm_cache import warm_dfs_cache\n", | ||
| "\n", | ||
| "source = RelBenchDatasetTask(DATASET, TASK)\n", | ||
| "assert source.task.task_type is TaskType.REGRESSION\n", | ||
| "# Default depth on purpose: the matrix is cached under its max_depth, and\n", | ||
| "# Nori-Rel slices the shared deepest matrix down to depth 2.\n", | ||
| "warm_dfs_cache(source, CacheConfig(FEATURE_CACHE, on_miss=\"fill\"))\n", | ||
| "print(f\"DFS cache ready at {FEATURE_CACHE}\")" | ||
| ] | ||
| }, | ||
| { | ||
| "cell_type": "markdown", | ||
| "id": "7cdc8c89c7104fffa095e18ddfef8986", | ||
| "metadata": {}, | ||
| "source": [ | ||
| "## Step 5 — run Nori-Rel\n", | ||
| "\n", | ||
| "Importing `relarena.models` registers the built-in models. `run_experiment` then runs the nested temporal protocol: fit on train and score validation, select the fixed configuration, refit on train + validation, and evaluate once on the hidden test split.\n", | ||
| "\n", | ||
| "The checkpoint downloads and its SHA-256 is verified on the first fit. Run only one Nori-Rel task per GPU process." | ||
| ] | ||
| }, | ||
| { | ||
| "cell_type": "code", | ||
| "execution_count": null, | ||
| "id": "b118ea5561624da68c537baed56e602f", | ||
| "metadata": {}, | ||
| "outputs": [], | ||
| "source": [ | ||
| "import relarena.models # noqa: F401 — registers Nori-Rel and the other models\n", | ||
| "from relarena import registry, run_experiment\n", | ||
| "\n", | ||
| "model_cls = registry.get(\"nori-rel\")\n", | ||
| "summary = run_experiment(\n", | ||
| " model_cls,\n", | ||
| " DATASET,\n", | ||
| " TASK,\n", | ||
| " seed=SEED,\n", | ||
| " n_trials=N_TRIALS,\n", | ||
| " cache_dir=FEATURE_CACHE,\n", | ||
| " cache_predictions=True,\n", | ||
| ")\n", | ||
| "summary" | ||
| ] | ||
| }, | ||
| { | ||
| "cell_type": "markdown", | ||
| "id": "938c804e27f84196a10c8828c723f798", | ||
| "metadata": {}, | ||
| "source": [ | ||
| "## Step 6 — inspect the result\n", | ||
| "\n", | ||
| "The result table contains every evaluated configuration, validation and test metrics, and separate fit/predict timings. Nori-Rel has one fixed configuration, so the single row is both the default and selected result. For MAE, lower is better." | ||
| ] | ||
| }, | ||
| { | ||
| "cell_type": "code", | ||
| "execution_count": null, | ||
| "id": "504fb2a444614c0babb325280ed9130a", | ||
| "metadata": {}, | ||
| "outputs": [], | ||
| "source": [ | ||
| "from relarena import summary_to_dataframe\n", | ||
| "\n", | ||
| "results = summary_to_dataframe(summary)\n", | ||
| "headline_columns = [\n", | ||
| " column\n", | ||
| " for column in (\n", | ||
| " \"model\",\n", | ||
| " \"dataset\",\n", | ||
| " \"task\",\n", | ||
| " \"metric\",\n", | ||
| " \"selected\",\n", | ||
| " \"val_score\",\n", | ||
| " \"test_score\",\n", | ||
| " \"fit_time_tuning\",\n", | ||
| " \"predict_time_tuning\",\n", | ||
| " \"fit_time_refit\",\n", | ||
| " \"predict_time_refit\",\n", | ||
| " )\n", | ||
| " if column in results.columns\n", | ||
| "]\n", | ||
| "results[headline_columns]" | ||
| ] | ||
| }, | ||
| { | ||
| "cell_type": "markdown", | ||
| "id": "59bbdb311c014d738909a11f9e486628", | ||
| "metadata": {}, | ||
| "source": [ | ||
| "## Step 7 — save the reproducible CSV\n", | ||
| "\n", | ||
| "The CSV follows RelArena's shared results schema and can be concatenated with other task runs before leaderboard analysis." | ||
| ] | ||
| }, | ||
| { | ||
| "cell_type": "code", | ||
| "execution_count": null, | ||
| "id": "b43b363d81ae4b689946ece5c682cd59", | ||
| "metadata": {}, | ||
| "outputs": [], | ||
| "source": [ | ||
| "output_path = RESULTS_DIR / f\"{DATASET}-{TASK}-seed-{SEED}.csv\"\n", | ||
| "results.to_csv(output_path, index=False)\n", | ||
| "print(f\"Saved {output_path.resolve()}\")" | ||
| ] | ||
| }, | ||
| { | ||
| "cell_type": "markdown", | ||
| "id": "8a65eabff63a45729fe45fb5ade58bdc", | ||
| "metadata": {}, | ||
| "source": [ | ||
| "## Command-line equivalent\n", | ||
| "\n", | ||
| "After choosing the same paths, these commands perform the cache warm and experiment without Jupyter:\n", | ||
| "\n", | ||
| "```bash\n", | ||
| "uv run --all-packages --no-sync python -m relarena.featurization.warm_cache \\\n", | ||
| " --dataset rel-f1 --task driver-position \\\n", | ||
| " --cache-dir ~/.cache/relarena/nori-rel/features\n", | ||
| "\n", | ||
| "CUDA_VISIBLE_DEVICES=0 uv run --all-packages --no-sync relarena \\\n", | ||
| " --model nori-rel --datasets rel-f1 --tasks driver-position \\\n", | ||
| " --seed 0 --n-trials 1 \\\n", | ||
| " --cache-dir ~/.cache/relarena/nori-rel/features \\\n", | ||
| " --output results/nori-rel/rel-f1-driver-position-seed-0.csv\n", | ||
| "```\n", | ||
| "\n", | ||
| "### Troubleshooting\n", | ||
| "\n", | ||
| "- **Unsupported task type:** choose a row from the regression task table in Step 3; classification is intentionally unsupported.\n", | ||
| "- **Cache miss:** rerun Step 4 with the same `DATASET`, `TASK`, and `FEATURE_CACHE`.\n", | ||
| "- **CUDA out of memory:** do not run concurrent tasks on one GPU. The adapter already bounds each forward pass and uses a host-offloaded BF16 cache without lossy fallbacks.\n", | ||
| "- **Slow first run:** DFS is CPU-heavy and the initial run downloads data and checkpoint files. Reusing `FEATURE_CACHE`, `HF_HOME`, and `RELBENCH_CACHE_DIR` avoids repeating that work." | ||
| ] | ||
| } | ||
| ], | ||
| "metadata": { | ||
| "kernelspec": { | ||
| "display_name": "Python 3 (ipykernel)", | ||
| "language": "python", | ||
| "name": "python3" | ||
| }, | ||
| "language_info": { | ||
| "name": "python", | ||
| "version": "3.11" | ||
| } | ||
| }, | ||
| "nbformat": 4, | ||
| "nbformat_minor": 5 | ||
| } | ||
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -0,0 +1,5 @@ | ||
| """Nori-Rel: depth-2 DFS features with the frozen Nori 30M regressor.""" | ||
|
Collaborator
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. I know that you describe how to add warm the cache in the notebook, but could you also add a |
||
|
|
||
| from relarena.models.nori_rel.model import NoriRelModel | ||
|
|
||
| __all__ = ["NoriRelModel"] | ||
Oops, something went wrong.
Add this suggestion to a batch that can be applied as a single commit.
This suggestion is invalid because no changes were made to the code.
Suggestions cannot be applied while the pull request is closed.
Suggestions cannot be applied while viewing a subset of changes.
Only one suggestion per line can be applied in a batch.
Add this suggestion to a batch that can be applied as a single commit.
Applying suggestions on deleted lines is not supported.
You must change the existing code in this line in order to create a valid suggestion.
Outdated suggestions cannot be applied.
This suggestion has been applied or marked resolved.
Suggestions cannot be applied from pending reviews.
Suggestions cannot be applied on multi-line comments.
Suggestions cannot be applied while the pull request is queued to merge.
Suggestion cannot be applied right now. Please check back later.
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
Optional suggestion: You might want to add cell outputs, so users can easily read the notebook without having to run all cells themselves.
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
Will re-run it end to end and commit the output