Spaces:
Sleeping
Sleeping
| """Prepare a MemisisLabs dataset for TabDiff (github.com/MinkaiXu/TabDiff). | |
| Writes the CSV + Info JSON into a TabDiff checkout's data/ layout so you can run | |
| `process_dataset.py`, then train/sample. Column-type indices (num/cat/target) are | |
| auto-detected from our dataset registry. | |
| Usage (on the lab server, inside the TabDiff repo's conda env): | |
| python scripts/tabdiff_prepare.py --dataset openml_45040 \ | |
| --name schizophrenia --tabdiff-root /path/to/TabDiff | |
| """ | |
| from __future__ import annotations | |
| import argparse | |
| import json | |
| import sys | |
| from pathlib import Path | |
| sys.path.insert(0, str(Path(__file__).resolve().parent.parent)) | |
| from pipeline import datasets # noqa: E402 | |
| from pipeline.metadata import is_categorical # noqa: E402 | |
| def main() -> None: | |
| ap = argparse.ArgumentParser(description=__doc__) | |
| ap.add_argument("--dataset", default="openml_45040", help="MemisisLabs dataset id") | |
| ap.add_argument("--name", default="schizophrenia", help="TabDiff dataset name") | |
| ap.add_argument("--tabdiff-root", required=True, help="path to the cloned TabDiff repo") | |
| args = ap.parse_args() | |
| real, target, _protected, task = datasets.load(args.dataset) | |
| cols = list(real.columns) | |
| target_idx = cols.index(target) | |
| num_idx, cat_idx = [], [] | |
| for i, c in enumerate(cols): | |
| if i == target_idx: | |
| continue | |
| (cat_idx if is_categorical(real[c]) else num_idx).append(i) | |
| root = Path(args.tabdiff_root) | |
| (root / "data" / args.name).mkdir(parents=True, exist_ok=True) | |
| (root / "data" / "Info").mkdir(parents=True, exist_ok=True) | |
| csv_path = root / "data" / args.name / f"{args.name}.csv" | |
| real.to_csv(csv_path, index=False) | |
| info = { | |
| "name": args.name, | |
| "task_type": "binclass" if task == "binclass" else "regression", | |
| "header": "infer", | |
| "column_names": None, | |
| "num_col_idx": num_idx, | |
| "cat_col_idx": cat_idx, | |
| "target_col_idx": [target_idx], | |
| "file_type": "csv", | |
| "data_path": f"data/{args.name}/{args.name}.csv", | |
| "test_path": None, | |
| } | |
| info_path = root / "data" / "Info" / f"{args.name}.json" | |
| info_path.write_text(json.dumps(info, indent=4)) | |
| print(f"wrote {csv_path} ({real.shape})") | |
| print(f"wrote {info_path}") | |
| print(f" num_col_idx={num_idx}\n cat_col_idx={cat_idx}\n target_col_idx=[{target_idx}] ({target})") | |
| print("\nNext (in the TabDiff repo):") | |
| print(f" python process_dataset.py --dataname {args.name}") | |
| print(f" python main.py --dataname {args.name} --mode train --no_wandb") | |
| print(f" python main.py --dataname {args.name} --mode test --report --no_wandb") | |
| if __name__ == "__main__": | |
| main() | |