-
Notifications
You must be signed in to change notification settings - Fork 1
Expand file tree
/
Copy pathinference_collision_2objects.py
More file actions
101 lines (82 loc) · 3.49 KB
/
Copy pathinference_collision_2objects.py
File metadata and controls
101 lines (82 loc) · 3.49 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
import os
# NCCL/UCX settings must be applied before importing torch or jax.
os.environ.setdefault("NCCL_IB_DISABLE", "1")
os.environ.setdefault("NCCL_P2P_DISABLE", "1")
os.environ.setdefault("NCCL_SOCKET_IFNAME", "^docker0,lo")
os.environ.setdefault("UCX_TLS", "tcp,sockcm")
os.environ.setdefault("UCX_NET_DEVICES", "")
import hydra
import omegaconf
import wandb
from omegaconf import DictConfig, open_dict
from experiments.fitting import get_model_pde
from experiments.fitting.datasets import get_inference_dataloader
from experiments.fitting.trainers.pde_infer_2objects_pair_vel_encoder import EncoderDecoder
from experiments.fitting.trainers.trainer_utils import solvers_infer
@hydra.main(version_base=None, config_path="./config/", config_name="config_collision_2objects_vel_inference")
def inference(cfg: DictConfig):
solvers_infer.COLLISION_RADIUS = cfg.encoder.collision_radius
solvers_infer.COLLISION_DETECT_RADIUS = cfg.encoder.collision_detect_radius
solvers_infer.OBJECT_NUM = cfg.inference.object_num
with open_dict(cfg):
cfg.inference.use = True
cfg.loading.load_checkpoint = True
if not cfg.logging.log_dir:
hydra_cfg = hydra.core.hydra_config.HydraConfig.get()
cfg.logging.log_dir = hydra_cfg["runtime"]["output_dir"]
wandb.init(
project=f"{cfg.proj_name}-inference",
dir=cfg.logging.log_dir,
config=omegaconf.OmegaConf.to_container(cfg),
mode="disabled",
)
testset_geo, testset_combo, testset_multi = get_inference_dataloader(cfg=cfg)
with open_dict(cfg.dataset):
cfg.dataset.image_shape = next(iter(testset_geo))[0][0][0].shape
with open_dict(cfg.encoder):
cfg.encoder.use = True
nef, encoder, ode_model = get_model_pde(cfg)
trainer = EncoderDecoder(
nef=nef,
pointnet=encoder,
ode_model=ode_model,
config=cfg,
train_loader=[testset_geo],
val_loader=[testset_combo],
vis_loader=[testset_multi],
seed=cfg.seed,
)
init_state = trainer.init_train_state()
loaded_state = trainer.load_checkpoint(path=cfg.loading.load_path)
state = init_state.replace(
params={
"nef": loaded_state.params["nef"],
"encoder": loaded_state.params["encoder"],
"ode_params": loaded_state.params["ode_params"],
},
nef_opt_state=loaded_state.nef_opt_state,
encoder_opt_state=loaded_state.encoder_opt_state,
ode_opt_state=loaded_state.ode_opt_state,
)
print(f"Loaded checkpoint from {cfg.loading.load_path}")
if cfg.inference.get("run_geo", True):
for batch_idx, batch in enumerate(testset_geo):
trainer.visualize_batch_inference(
state, batch, name=f"geo_inference{batch_idx}"
)
if cfg.inference.get("run_comb", True):
for batch_idx, batch in enumerate(testset_combo):
trainer.visualize_batch_inference(
state, batch, name=f"comb_inference{batch_idx}"
)
if cfg.inference.get("run_multi", False):
for batch_idx, batch in enumerate(testset_multi):
trainer.visualize_batch_multi_inference(
state, batch, name=f"multi_inference{batch_idx}"
)
print(f"Inference finished. Results saved to {cfg.logging.log_dir}")
if __name__ == "__main__":
if "CUDA_VISIBLE_DEVICES" not in os.environ:
os.environ["CUDA_VISIBLE_DEVICES"] = "0"
os.environ.setdefault("XLA_PYTHON_CLIENT_MEM_FRACTION", "0.95")
inference()