AUREOLE-R-v3 / scripts /train_prior.py
PureOne's picture
AUREOLE-R 3.0.0-hf.1: standalone public research release
9d6c005 verified
Raw
History Blame Contribute Delete
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)