MCPcopy Create free account
hub / github.com/YesianRohn/TextSSR / log_validation

Function log_validation

diffusers/examples/controlnet/train_controlnet_sd3.py:78–179  ·  view source on GitHub ↗
(controlnet, args, accelerator, weight_dtype, step, is_final_validation=False)

Source from the content-addressed store, hash-verified

76
77
78def log_validation(controlnet, args, accelerator, weight_dtype, step, is_final_validation=False):
79 logger.info("Running validation... ")
80
81 if not is_final_validation:
82 controlnet = accelerator.unwrap_model(controlnet)
83 else:
84 controlnet = SD3ControlNetModel.from_pretrained(args.output_dir, torch_dtype=weight_dtype)
85
86 pipeline = StableDiffusion3ControlNetPipeline.from_pretrained(
87 args.pretrained_model_name_or_path,
88 controlnet=controlnet,
89 safety_checker=None,
90 revision=args.revision,
91 variant=args.variant,
92 torch_dtype=weight_dtype,
93 )
94 pipeline = pipeline.to(torch.device(accelerator.device))
95 pipeline.set_progress_bar_config(disable=True)
96
97 if args.seed is None:
98 generator = None
99 else:
100 generator = torch.manual_seed(args.seed)
101
102 if len(args.validation_image) == len(args.validation_prompt):
103 validation_images = args.validation_image
104 validation_prompts = args.validation_prompt
105 elif len(args.validation_image) == 1:
106 validation_images = args.validation_image * len(args.validation_prompt)
107 validation_prompts = args.validation_prompt
108 elif len(args.validation_prompt) == 1:
109 validation_images = args.validation_image
110 validation_prompts = args.validation_prompt * len(args.validation_image)
111 else:
112 raise ValueError(
113 "number of `args.validation_image` and `args.validation_prompt` should be checked in `parse_args`"
114 )
115
116 image_logs = []
117 inference_ctx = contextlib.nullcontext() if is_final_validation else torch.autocast(accelerator.device.type)
118
119 for validation_prompt, validation_image in zip(validation_prompts, validation_images):
120 validation_image = Image.open(validation_image).convert("RGB")
121
122 images = []
123
124 for _ in range(args.num_validation_images):
125 with inference_ctx:
126 image = pipeline(
127 validation_prompt, control_image=validation_image, num_inference_steps=20, generator=generator
128 ).images[0]
129
130 images.append(image)
131
132 image_logs.append(
133 {"validation_image": validation_image, "images": images, "validation_prompt": validation_prompt}
134 )
135

Callers 1

mainFunction · 0.70

Calls 6

free_memoryFunction · 0.90
infoMethod · 0.80
from_pretrainedMethod · 0.45
toMethod · 0.45
deviceMethod · 0.45

Tested by

no test coverage detected