| """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() |
|
|