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

Function log_validation

diffusers/examples/controlnet/train_controlnet_flux.py:75–198  ·  view source on GitHub ↗
(
    vae, flux_transformer, flux_controlnet, args, accelerator, weight_dtype, step, is_final_validation=False
)

Source from the content-addressed store, hash-verified

73
74
75def log_validation(
76 vae, flux_transformer, flux_controlnet, args, accelerator, weight_dtype, step, is_final_validation=False
77):
78 logger.info("Running validation... ")
79
80 if not is_final_validation:
81 flux_controlnet = accelerator.unwrap_model(flux_controlnet)
82 pipeline = FluxControlNetPipeline.from_pretrained(
83 args.pretrained_model_name_or_path,
84 controlnet=flux_controlnet,
85 transformer=flux_transformer,
86 torch_dtype=torch.bfloat16,
87 )
88 else:
89 flux_controlnet = FluxControlNetModel.from_pretrained(
90 args.output_dir, torch_dtype=torch.bfloat16, variant=args.save_weight_dtype
91 )
92 pipeline = FluxControlNetPipeline.from_pretrained(
93 args.pretrained_model_name_or_path,
94 controlnet=flux_controlnet,
95 transformer=flux_transformer,
96 torch_dtype=torch.bfloat16,
97 )
98
99 pipeline.to(accelerator.device)
100 pipeline.set_progress_bar_config(disable=True)
101
102 if args.enable_xformers_memory_efficient_attention:
103 pipeline.enable_xformers_memory_efficient_attention()
104
105 if args.seed is None:
106 generator = None
107 else:
108 generator = torch.Generator(device=accelerator.device).manual_seed(args.seed)
109
110 if len(args.validation_image) == len(args.validation_prompt):
111 validation_images = args.validation_image
112 validation_prompts = args.validation_prompt
113 elif len(args.validation_image) == 1:
114 validation_images = args.validation_image * len(args.validation_prompt)
115 validation_prompts = args.validation_prompt
116 elif len(args.validation_prompt) == 1:
117 validation_images = args.validation_image
118 validation_prompts = args.validation_prompt * len(args.validation_image)
119 else:
120 raise ValueError(
121 "number of `args.validation_image` and `args.validation_prompt` should be checked in `parse_args`"
122 )
123
124 image_logs = []
125 if is_final_validation or torch.backends.mps.is_available():
126 autocast_ctx = nullcontext()
127 else:
128 autocast_ctx = torch.autocast(accelerator.device.type)
129
130 for validation_prompt, validation_image in zip(validation_prompts, validation_images):
131 from diffusers.utils import load_image
132

Callers 1

mainFunction · 0.70

Calls 9

load_imageFunction · 0.90
free_memoryFunction · 0.90
infoMethod · 0.80
from_pretrainedMethod · 0.45
toMethod · 0.45
resizeMethod · 0.45
encode_promptMethod · 0.45

Tested by

no test coverage detected