MCPcopy Create free account
hub / github.com/TeleHuman/TextOp / load_dar

Method load_dar

TextOpDeploy/src/textop_ctrl/scripts/rmdar.py:344–422  ·  view source on GitHub ↗
(self)

Source from the content-addressed store, hash-verified

342 f"in {(t1 - t0) * 1000:.2f} ms")
343
344 def load_dar(self):
345 ckpt = Path(self.ckpt)
346 dar_cfg_path = ckpt.parent / '.hydra' / 'config.yaml'
347 # dar_cfg_path = ckpt.parent / 'cfg.yaml'
348 self.dar_cfg = OmegaConf.load(dar_cfg_path)
349
350 cfg = self.dar_cfg
351 cfg.device = str(self._device)
352 cfg.ckpt.dar = str(ckpt)
353 cfg.train.manager.device = str(self._device)
354 cfg.train.manager.platform._target_ = 'robotmdar.train.train_platforms.NoPlatform'
355 cfg.data.datadir = self.config.dar.datadir
356 cfg.skeleton.asset.assetRoot = self.config.dar.skeleton_assetRoot
357 cfg.data.val.split = 'none'
358 cfg.data.val.batch_size = 1
359 cfg.use_full_sample = self.use_full_sample
360 cfg.guidance_scale = self.guidance_scale
361
362 seed.set(cfg.seed)
363 self.clip_model = load_and_freeze_clip("ViT-B/32", device=self._device)
364 val_data: Dataset = instantiate(cfg.data.val)
365 vae: VAE = instantiate(cfg.vae)
366 denoiser: Denoiser = instantiate(cfg.denoiser)
367
368 schedule_sampler: SSampler = instantiate(
369 cfg.diffusion.schedule_sampler)
370 diffusion: Diffusion = schedule_sampler.diffusion
371
372 vae.eval()
373 denoiser.eval()
374
375 # Load checkpoints
376 manager: DARManager = instantiate(cfg.train.manager)
377 manager.hold_model(vae, denoiser, None, val_data)
378
379 # vae_trt = vae
380 # denoiser_trt = denoiser
381 try:
382 vae_trt = torch.compile(vae, backend='tensorrt')
383 denoiser_trt = torch.compile(denoiser, backend='tensorrt')
384 except KeyError as e:
385 error_key = e.args[0]
386 if error_key == 'torch_dynamo_backends':
387 # Now we are in jetson env, which does not support torch.compile
388 vae_trt, denoiser_trt = jetson_compatible_torch_compile(
389 vae, denoiser, cfg)
390 else:
391 raise e
392
393 # vae_trt = torch.compile(vae, backend='inductor')
394 # denoiser_trt = torch.compile(denoiser, backend='inductor')
395
396 cfg_denoiser = ClassifierFreeWrapper(denoiser_trt)
397
398 self.dataset = val_data
399 self.dataiter = iter(val_data)
400 self._motion_dt = 1 / val_data.fps
401 self.future_len = cfg.data.future_len

Callers 1

__init__Method · 0.95

Calls 10

warmupMethod · 0.95
_reset_motion_bufferMethod · 0.95
load_and_freeze_clipFunction · 0.90
loadMethod · 0.80
setMethod · 0.80
evalMethod · 0.80
hold_modelMethod · 0.45

Tested by

no test coverage detected