Download scripts/train.py from OneScience-Group/MeteoNorm-RF: direct link, hf CLI and curl.
- Browser
- Download file 3.33 kB
-
https://huggingface.co/OneScience-Group/MeteoNorm-RF/resolve/main/scripts/train.py
- Command line
-
hf download hf://OneScience-Group/MeteoNorm-RF/scripts/train.py
-
curl -L -o train.py https://huggingface.co/OneScience-Group/MeteoNorm-RF/resolve/main/scripts/train.py
3.33 kB
| """Train a pure NumPy RF; torchrun ranks build disjoint tree subsets.""" | |
| import json | |
| import os | |
| import pickle | |
| import sys | |
| from pathlib import Path | |
| import numpy as np | |
| import yaml | |
| ROOT = Path(__file__).resolve().parents[1] | |
| sys.path.insert(0, str(ROOT)) | |
| from model.meteonorm_rf import (FEATURE_NAMES, FORMAT_VERSION, MODEL_NAME, | |
| MultiOutputRandomForest, MeteoNormRF, | |
| encode_features, merge_states, save_checkpoint) | |
| def main(): | |
| config = yaml.safe_load((ROOT / "conf/config.yaml").read_text()) | |
| seed = int(config["seed"]) | |
| data = np.load(ROOT / config["data"]["path"]) | |
| x, y = encode_features(data), data["pollution"].astype(np.float32) | |
| rng = np.random.default_rng(seed) | |
| order = rng.permutation(len(x)) | |
| cut = int(float(config["data"]["train_fraction"]) * len(order)) | |
| train_indices, test_indices = order[:cut], order[cut:] | |
| world, rank = int(os.environ.get("WORLD_SIZE", "1")), int(os.environ.get("RANK", "0")) | |
| options = config["model"]["engineering"] | |
| trees = int(options["trees"]) | |
| assigned = list(range(rank, trees, world)) | |
| forest = MultiOutputRandomForest(n_trees=trees, seed=seed, | |
| max_depth=int(options["max_depth"]), | |
| min_samples_leaf=int(options["min_samples_leaf"]), | |
| max_features=options["max_features"], | |
| split_candidates=int(options["split_candidates"])) | |
| model = MeteoNormRF(forest).fit(x[train_indices], y[train_indices], assigned) | |
| checkpoint = ROOT / config["paths"]["checkpoint"] | |
| checkpoint.parent.mkdir(parents=True, exist_ok=True) | |
| shard = checkpoint.with_suffix(f".rank{rank}.pkl") | |
| with open(shard, "wb") as stream: | |
| pickle.dump(model.state_dict(), stream) | |
| if world > 1: | |
| import torch | |
| torch.distributed.init_process_group("gloo") | |
| torch.distributed.barrier() | |
| if rank == 0: | |
| states = [] | |
| for item in range(world): | |
| with open(checkpoint.with_suffix(f".rank{item}.pkl"), "rb") as stream: | |
| states.append(pickle.load(stream)) | |
| merged = MeteoNormRF.from_state_dict(merge_states(states)) | |
| metadata = {"feature_names": FEATURE_NAMES, "train_indices": train_indices, | |
| "test_indices": test_indices, "split": "seeded random 70/30", | |
| "paper_trees": config["paper_protocol"]["trees"], "engineering_trees": trees} | |
| save_checkpoint(checkpoint, merged, metadata) | |
| metrics = ROOT / config["paths"]["training_metrics"] | |
| metrics.parent.mkdir(parents=True, exist_ok=True) | |
| metrics.write_text(json.dumps({"format_version": FORMAT_VERSION, "model": MODEL_NAME, | |
| "rows": len(x), "train_rows": cut, "test_rows": len(x) - cut, | |
| "trees": trees, "world_size": world}, indent=2) + "\n") | |
| for item in range(world): | |
| checkpoint.with_suffix(f".rank{item}.pkl").unlink() | |
| print(f"checkpoint={checkpoint.relative_to(ROOT)} trees={trees} train={cut} test={len(x)-cut}") | |
| if world > 1: | |
| torch.distributed.barrier() | |
| torch.distributed.destroy_process_group() | |
| if __name__ == "__main__": | |
| main() | |