MCPcopy Create free account
hub / github.com/CompVis/zigma / get_data_generator

Function get_data_generator

sample_acc.py:278–315  ·  view source on GitHub ↗
()

Source from the content-addressed store, hash-verified

276 total = 0
277
278 def get_data_generator():
279 while True:
280 for data in tqdm(
281 loader,
282 disable=not (rank == 0),
283 initial=0,
284 desc=f"generate_images, for iters {iterations}",
285 ):
286 if args.use_latent:
287 if has_text(args):
288 _cap_feats = data["caption_feature"]
289 B, N, T, C = _cap_feats.shape # each image has N captions
290 yield data["img_feature"].to(device), _cap_feats[
291 :, random.randint(0, N - 1)
292 ].to(device)
293 elif "facehq" in str(args.data.name):
294 yield data["latent"].to(device), None
295 elif "church" in str(args.data.name):
296 yield data["latent"].to(device), None
297 elif "ucf101" in str(args.data.name):
298 yield data["frame_feature256"].to(device), data["cls_id"]
299 elif "celebav" in str(args.data.name):
300 _start = random.randint(
301 0,
302 data["frame_feature256"].shape[1]
303 - args.model.params.video_frames
304 - 1,
305 )
306 _video = data["frame_feature256"][
307 :, _start : _start + args.model.params.video_frames
308 ].to(device)
309 yield _video, None
310 else:
311 raise NotImplementedError(
312 f"latent data not supported, args.data.name={args.data.name}"
313 )
314 else:
315 yield data["image"].to(device), None
316
317 data_generator = get_data_generator()
318

Callers 1

mainFunction · 0.70

Calls 2

has_textFunction · 0.90
toMethod · 0.80

Tested by

no test coverage detected