MCPcopy Create free account
hub / github.com/GasaiYU/PartRM / prepare_batch_data

Method prepare_batch_data

preprocess/src/model.py:42–102  ·  view source on GitHub ↗
(self, batch)

Source from the content-addressed store, hash-verified

40 os.makedirs(os.path.join(self.logdir, 'images_val'), exist_ok=True)
41
42 def prepare_batch_data(self, batch):
43 lrm_generator_input = {}
44 render_gt = {} # for supervision
45
46 # input images
47 images = batch['input_images']
48 images = v2.functional.resize(
49 images, self.input_size, interpolation=3, antialias=True).clamp(0, 1)
50
51 lrm_generator_input['images'] = images.to(self.device)
52
53 # input cameras and render cameras
54 input_c2ws = batch['input_c2ws'].flatten(-2)
55 input_Ks = batch['input_Ks'].flatten(-2)
56 target_c2ws = batch['target_c2ws'].flatten(-2)
57 target_Ks = batch['target_Ks'].flatten(-2)
58 render_cameras_input = torch.cat([input_c2ws, input_Ks], dim=-1)
59 render_cameras_target = torch.cat([target_c2ws, target_Ks], dim=-1)
60 render_cameras = torch.cat([render_cameras_input, render_cameras_target], dim=1)
61
62 input_extrinsics = input_c2ws[:, :, :12]
63 input_intrinsics = torch.stack([
64 input_Ks[:, :, 0], input_Ks[:, :, 4],
65 input_Ks[:, :, 2], input_Ks[:, :, 5],
66 ], dim=-1)
67 cameras = torch.cat([input_extrinsics, input_intrinsics], dim=-1)
68
69 # add noise to input cameras
70 cameras = cameras + torch.rand_like(cameras) * 0.04 - 0.02
71
72 lrm_generator_input['cameras'] = cameras.to(self.device)
73 lrm_generator_input['render_cameras'] = render_cameras.to(self.device)
74
75 # target images
76 target_images = torch.cat([batch['input_images'], batch['target_images']], dim=1)
77 target_depths = torch.cat([batch['input_depths'], batch['target_depths']], dim=1)
78 target_alphas = torch.cat([batch['input_alphas'], batch['target_alphas']], dim=1)
79
80 # random crop
81 render_size = np.random.randint(self.render_size, 513)
82 target_images = v2.functional.resize(
83 target_images, render_size, interpolation=3, antialias=True).clamp(0, 1)
84 target_depths = v2.functional.resize(
85 target_depths, render_size, interpolation=0, antialias=True)
86 target_alphas = v2.functional.resize(
87 target_alphas, render_size, interpolation=0, antialias=True)
88
89 crop_params = v2.RandomCrop.get_params(
90 target_images, output_size=(self.render_size, self.render_size))
91 target_images = v2.functional.crop(target_images, *crop_params)
92 target_depths = v2.functional.crop(target_depths, *crop_params)[:, :, 0:1]
93 target_alphas = v2.functional.crop(target_alphas, *crop_params)[:, :, 0:1]
94
95 lrm_generator_input['render_size'] = render_size
96 lrm_generator_input['crop_params'] = crop_params
97
98 render_gt['target_images'] = target_images.to(self.device)
99 render_gt['target_depths'] = target_depths.to(self.device)

Callers 1

training_stepMethod · 0.95

Calls

no outgoing calls

Tested by

no test coverage detected