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

Method prepare_batch_data

preprocess/src/model_mesh.py:80–137  ·  view source on GitHub ↗
(self, batch)

Source from the content-addressed store, hash-verified

78 os.makedirs(os.path.join(self.logdir, 'images_val'), exist_ok=True)
79
80 def prepare_batch_data(self, batch):
81 lrm_generator_input = {}
82 render_gt = {}
83
84 # input images
85 images = batch['input_images']
86 images = v2.functional.resize(
87 images, self.input_size, interpolation=3, antialias=True).clamp(0, 1)
88
89 lrm_generator_input['images'] = images.to(self.device)
90
91 # input cameras and render cameras
92 input_c2ws = batch['input_c2ws']
93 input_Ks = batch['input_Ks']
94 target_c2ws = batch['target_c2ws']
95
96 render_c2ws = torch.cat([input_c2ws, target_c2ws], dim=1)
97 render_w2cs = torch.linalg.inv(render_c2ws)
98
99 input_extrinsics = input_c2ws.flatten(-2)
100 input_extrinsics = input_extrinsics[:, :, :12]
101 input_intrinsics = input_Ks.flatten(-2)
102 input_intrinsics = torch.stack([
103 input_intrinsics[:, :, 0], input_intrinsics[:, :, 4],
104 input_intrinsics[:, :, 2], input_intrinsics[:, :, 5],
105 ], dim=-1)
106 cameras = torch.cat([input_extrinsics, input_intrinsics], dim=-1)
107
108 # add noise to input_cameras
109 cameras = cameras + torch.rand_like(cameras) * 0.04 - 0.02
110
111 lrm_generator_input['cameras'] = cameras.to(self.device)
112 lrm_generator_input['render_cameras'] = render_w2cs.to(self.device)
113
114 # target images
115 target_images = torch.cat([batch['input_images'], batch['target_images']], dim=1)
116 target_depths = torch.cat([batch['input_depths'], batch['target_depths']], dim=1)
117 target_alphas = torch.cat([batch['input_alphas'], batch['target_alphas']], dim=1)
118 target_normals = torch.cat([batch['input_normals'], batch['target_normals']], dim=1)
119
120 render_size = self.render_size
121 target_images = v2.functional.resize(
122 target_images, render_size, interpolation=3, antialias=True).clamp(0, 1)
123 target_depths = v2.functional.resize(
124 target_depths, render_size, interpolation=0, antialias=True)
125 target_alphas = v2.functional.resize(
126 target_alphas, render_size, interpolation=0, antialias=True)
127 target_normals = v2.functional.resize(
128 target_normals, render_size, interpolation=3, antialias=True)
129
130 lrm_generator_input['render_size'] = render_size
131
132 render_gt['target_images'] = target_images.to(self.device)
133 render_gt['target_depths'] = target_depths.to(self.device)
134 render_gt['target_alphas'] = target_alphas.to(self.device)
135 render_gt['target_normals'] = target_normals.to(self.device)
136
137 return lrm_generator_input, render_gt

Callers 1

training_stepMethod · 0.95

Calls

no outgoing calls

Tested by

no test coverage detected