Download scripts/train_prior.py from PureOne/AUREOLE-R-v3: direct link, hf CLI and curl.
- Browser
- Download file 4.73 kB
-
https://huggingface.co/PureOne/AUREOLE-R-v3/resolve/main/scripts/train_prior.py
- Command line
-
hf download hf://PureOne/AUREOLE-R-v3/scripts/train_prior.py
-
curl -L -o train_prior.py https://huggingface.co/PureOne/AUREOLE-R-v3/resolve/main/scripts/train_prior.py
4.73 kB
| """Train one small prior on procedural physical rays; held-out scene split.""" | |
| from pathlib import Path | |
| import argparse,json,time,platform,sys | |
| ROOT=Path(__file__).resolve().parents[1] | |
| sys.path.insert(0,str(ROOT)) | |
| import numpy as np | |
| from aureole.renderer import Scene, VisibilityPrior | |
| def dataset(seeds, per_scene): | |
| xx,yy=[],[] | |
| for seed in seeds: | |
| rng=np.random.default_rng(48000+seed) | |
| p=np.c_[rng.uniform(-1,1,(per_scene,2)),np.zeros(per_scene)] | |
| l=np.c_[rng.uniform(-0.85,0.85,(per_scene,2)),np.full(per_scene,2.2)] | |
| scene=Scene.create(seed) | |
| xx.append(scene.features(p,l)) | |
| yy.append(scene.visibility(p,l).astype(np.float32)) | |
| return np.concatenate(xx),np.concatenate(yy) | |
| def run(epochs=40,output="retrained"): | |
| import torch | |
| torch.set_num_threads(2) | |
| torch.manual_seed(20260919) | |
| np.random.seed(20260919) | |
| torch.use_deterministic_algorithms(True) | |
| start=time.perf_counter() | |
| x,y=dataset(range(48),1536) | |
| vx,vy=dataset(range(100,108),1536) | |
| tx,ty=dataset(range(200,208),1536) | |
| mean=x.mean(0); scale=np.maximum(x.std(0),1e-5) | |
| inputs=torch.from_numpy((x-mean)/scale); labels=torch.from_numpy(y[:,None]) | |
| valx=torch.from_numpy((vx-mean)/scale); valy=torch.from_numpy(vy[:,None]) | |
| model=torch.nn.Sequential(torch.nn.Linear(16,48),torch.nn.ReLU(),torch.nn.Linear(48,48),torch.nn.ReLU(),torch.nn.Linear(48,1)) | |
| optimizer=torch.optim.Adam(model.parameters(),lr=0.003) | |
| history=[]; best=float("inf"); best_state=None | |
| for epoch in range(epochs): | |
| order=torch.randperm(len(inputs)) | |
| for ids in order.split(2048): | |
| optimizer.zero_grad(set_to_none=True) | |
| logits=model(inputs[ids]); loss=torch.nn.functional.binary_cross_entropy_with_logits(logits,labels[ids]) | |
| loss.backward();optimizer.step() | |
| with torch.no_grad(): | |
| bce=float(torch.nn.functional.binary_cross_entropy_with_logits(model(valx),valy)) | |
| history.append({"epoch":epoch+1,"validation_bce":bce}) | |
| if bce<best: | |
| best=bce;best_state={k:v.detach().clone() for k,v in model.state_dict().items()} | |
| if (epoch+1)%10==0: | |
| print(f"epoch {epoch+1}/{epochs}; validation BCE {bce:.6f}",flush=True) | |
| model.load_state_dict(best_state) | |
| arrays={"mean":mean,"scale":scale} | |
| for i,j in enumerate((0,2,4)): | |
| arrays[f"w{i}"]=model[j].weight.detach().numpy() | |
| arrays[f"b{i}"]=model[j].bias.detach().numpy() | |
| destination=ROOT/output | |
| destination.joinpath("models").mkdir(parents=True,exist_ok=True) | |
| path=destination/"models/visibility_prior.npz" | |
| np.savez_compressed(path,**arrays) | |
| portable=VisibilityPrior(path) | |
| def metrics(features,targets): | |
| p=portable(features).astype(float) | |
| pc=np.clip(p,1e-7,1-1e-7) | |
| constant=float(y.mean()) | |
| return {"rays":len(targets),"visible_fraction":float(targets.mean()), | |
| "brier":float(np.mean((p-targets)**2)), | |
| "bce":float(np.mean(-targets*np.log(pc)-(1-targets)*np.log(1-pc))), | |
| "accuracy":float(np.mean((p>=0.5)==targets)), | |
| "constant_training_mean_brier":float(np.mean((constant-targets)**2))} | |
| with torch.no_grad(): | |
| torch_pred=torch.sigmoid(model(torch.from_numpy((tx-mean)/scale))).numpy().ravel() | |
| report={"seed":20260919,"architecture":[16,48,48,1],"parameters":sum(p.numel() for p in model.parameters()), | |
| "training_scene_ids":list(range(48)),"validation_scene_ids":list(range(100,108)),"test_scene_ids":list(range(200,208)), | |
| "epochs":epochs,"selected_epoch":int(np.argmin([v["validation_bce"] for v in history]))+1, | |
| "selection":"minimum validation BCE; test set not used for selection", "training":metrics(x,y), | |
| "validation":metrics(vx,vy),"test":metrics(tx,ty),"history":history, | |
| "numpy_torch_max_abs_error":float(np.max(np.abs(portable(tx)-torch_pred))), | |
| "elapsed_seconds":time.perf_counter()-start,"python":platform.python_version(),"torch":torch.__version__, | |
| "device":"cpu","claim_scope":"learned visibility prior for three-sphere direct-light scenes; not a learned unified renderer"} | |
| destination.joinpath("results").mkdir(exist_ok=True) | |
| (destination/"results/training.json").write_text(json.dumps(report,indent=2)+"\n") | |
| print(json.dumps({k:v for k,v in report.items() if k not in ("history","training_scene_ids")},indent=2)) | |
| if __name__=="__main__": | |
| p=argparse.ArgumentParser();p.add_argument("--epochs",type=int,default=40) | |
| p.add_argument("--output",default="retrained",help="Keep new weights separate from the bundled checkpoint") | |
| args=p.parse_args();run(args.epochs,args.output) | |