(prompts, choose_model, infer_mode, seed, n_samples, camera_args=None)
| 176 | |
| 177 | |
| 178 | def model_run(prompts, choose_model, infer_mode, seed, n_samples, camera_args=None): |
| 179 | traj_list = get_traj_list() |
| 180 | camera_dict = get_camera_dict() |
| 181 | |
| 182 | RT = process_camera(camera_dict, camera_args).reshape(-1,12) |
| 183 | traj_flow = process_traj(traj_list).transpose(3,0,1,2) |
| 184 | |
| 185 | if choose_model == BASE_MODEL[0]: |
| 186 | model = model_v1 |
| 187 | noise_shape = [1, 4, 16, 32, 32] |
| 188 | else: |
| 189 | model = model_v2 |
| 190 | noise_shape = [1, 4, 16, 40, 64] |
| 191 | unconditional_guidance_scale = 7.5 |
| 192 | unconditional_guidance_scale_temporal = None |
| 193 | |
| 194 | ddim_steps= 50 |
| 195 | ddim_eta=1.0 |
| 196 | cond_T=800 |
| 197 | |
| 198 | if n_samples < 1: |
| 199 | n_samples = 1 |
| 200 | if n_samples > 4: |
| 201 | n_samples = 4 |
| 202 | |
| 203 | seed_everything(seed) |
| 204 | |
| 205 | if infer_mode == MODE[0]: |
| 206 | camera_poses = RT |
| 207 | camera_poses = torch.tensor(camera_poses).float() |
| 208 | camera_poses = camera_poses.unsqueeze(0) |
| 209 | trajs = None |
| 210 | if torch.cuda.is_available(): |
| 211 | camera_poses = camera_poses.cuda() |
| 212 | elif infer_mode == MODE[1]: |
| 213 | trajs = traj_flow |
| 214 | trajs = torch.tensor(trajs).float() |
| 215 | trajs = trajs.unsqueeze(0) |
| 216 | camera_poses = None |
| 217 | if torch.cuda.is_available(): |
| 218 | trajs = trajs.cuda() |
| 219 | else: |
| 220 | camera_poses = RT |
| 221 | trajs = traj_flow |
| 222 | camera_poses = torch.tensor(camera_poses).float() |
| 223 | trajs = torch.tensor(trajs).float() |
| 224 | camera_poses = camera_poses.unsqueeze(0) |
| 225 | trajs = trajs.unsqueeze(0) |
| 226 | if torch.cuda.is_available(): |
| 227 | camera_poses = camera_poses.cuda() |
| 228 | trajs = trajs.cuda() |
| 229 | |
| 230 | |
| 231 | ddim_sampler = DDIMSampler(model) |
| 232 | batch_size = noise_shape[0] |
| 233 | ## get condition embeddings (support single prompt only) |
| 234 | if isinstance(prompts, str): |
| 235 | prompts = [prompts] |
nothing calls this directly
no test coverage detected