CrucibleAI/ControlNetMediaPipeFace
5761.2k
1import json2import os3import sys4from dataclasses import dataclass, field5from glob import glob6from typing import Mapping7 8from PIL import Image9from tqdm import tqdm10 11from laion_face_common import generate_annotation12 13 14@dataclass15class RunProgress:16 pending: list = field(default_factory=list)17 success: list = field(default_factory=list)18 skipped_size: list = field(default_factory=list)19 skipped_nsfw: list = field(default_factory=list)20 skipped_noface: list = field(default_factory=list)21 skipped_smallface: list = field(default_factory=list)22 23 24def main(25 status_filename: str,26 prompt_filename: str,27 input_glob: str,28 output_directory: str,29 annotated_output_directory: str = "",30 min_image_size: int = 384,31 max_image_size: int = 32766,32 min_face_size_pixels: int = 64,33 prompt_mapping: dict = None, # If present, maps a filename to a text prompt.34):35 status = RunProgress()36 37 if os.path.exists(status_filename):38 print("Continuing from checkpoint.")39 # Restore a saved state:40 status_temp = json.load(open(status_filename, 'rt'))41 for k in status.__dict__.keys():42 status.__setattr__(k, status_temp[k])43 # Output label file:44 pout = open(prompt_filename, 'at')45 else:46 print("Starting run.")47 status = RunProgress()48 status.pending = list(glob(input_glob))49 # Output label file:50 pout = open(prompt_filename, 'wt')51 with open(status_filename, 'wt') as fout:52 json.dump(status.__dict__, fout)53 54 print(f"{len(status.pending)} images remaining")55 56 # If we don't have a preexisting set of labels (like for ImageNet/MSCOCO), just null-fill the mapping.57 # We will try on a per-image basis to see if there's a metadata .json.58 if prompt_mapping is None:59 prompt_mapping = dict()60 61 step = 062 with tqdm(total=len(status.pending)) as pbar:63 while len(status.pending) > 0:64 full_filename = status.pending.pop()65 pbar.update(1)66 step += 167 68 if step % 100 == 0:69 # Checkpoint save:70 with open(status_filename, 'wt') as fout:71 json.dump(status.__dict__, fout)72 73 _fpath, fname = os.path.split(full_filename)74 75 # Make our output filenames.76 # We used to do this here so we could check if a file existed before writing, then skip it, but since we77 # have a 'status' that we cache and update, we no longer have to do this check.78 annotation_filename = ""79 if annotated_output_directory:80 annotation_filename = os.path.join(annotated_output_directory, fname)81 output_filename = os.path.join(output_directory, fname)82 83 # The LAION dataset has accompanying .json files with each image.84 partial_filename, extension = os.path.splitext(full_filename)85 candidate_json_fullpath = partial_filename + ".json"86 image_metadata = {}87 if os.path.exists(candidate_json_fullpath):88 try:89 image_metadata = json.load(open(candidate_json_fullpath, 'rt'))90 except Exception as e:91 print(e)92 if "NSFW" in image_metadata:93 nsfw_marker = image_metadata.get("NSFW") # This can be "", None, or other weird things.94 if nsfw_marker is not None and nsfw_marker.lower() != "unlikely":95 # Skip NSFW images.96 status.skipped_nsfw.append(full_filename)97 continue98 99 # Try to get a prompt/caption from the metadata or the prompt mapping.100 image_prompt = image_metadata.get("caption", prompt_mapping.get(fname, ""))101 102 # Load image:103 img = Image.open(full_filename).convert("RGB")104 img_width = img.size[0]105 img_height = img.size[1]106 img_size = min(img.size[0], img.size[1])107 if img_size < min_image_size or max(img_width, img_height) > max_image_size:108 status.skipped_size.append(full_filename)109 continue110 111 # We re-initialize the detector every time because it has a habit of triggering weird race conditions.112 empty, annotated, faces_before_filtering, faces_after_filtering = generate_annotation(113 img,114 max_faces=5,115 min_face_size_pixels=min_face_size_pixels,116 return_annotation_data=True117 )118 if faces_before_filtering == 0:119 # Skip images with no faces.120 status.skipped_noface.append(full_filename)121 continue122 if faces_after_filtering == 0:123 # Skip images with no faces large enough124 status.skipped_smallface.append(full_filename)125 continue126 127 Image.fromarray(empty).save(output_filename)128 if annotation_filename:129 Image.fromarray(annotated).save(annotation_filename)130 131 # See https://github.com/lllyasviel/ControlNet/blob/main/docs/train.md for the training file format.132 # prompt.json133 # a JSONL file with {"source": "source/0.jpg", "target": "target/0.jpg", "prompt": "..."}.134 # a source/xxxxx.jpg or source/xxxx.png file for each of the inputs.135 # a target/xxxxx.jpg for each of the outputs.136 pout.write(json.dumps({137 "source": os.path.join(output_directory, fname),138 "target": full_filename,139 "prompt": image_prompt,140 }) + "\n")141 pout.flush()142 status.success.append(full_filename)143 144 # We do save every 100 iterations, but it's good to save on completion, too.145 with open(status_filename, 'wt') as fout:146 json.dump(status.__dict__, fout)147 148 pout.close()149 print("Done!")150 print(f"{len(status.success)} images added to dataset.")151 print(f"{len(status.skipped_size)} images rejected for size.")152 print(f"{len(status.skipped_smallface)} images rejected for having faces too small.")153 print(f"{len(status.skipped_noface)} images rejected for not having faces.")154 print(f"{len(status.skipped_nsfw)} images rejected for NSFW.")155 156 157if __name__ == "__main__":158 if len(sys.argv) >= 3 and "-h" not in sys.argv:159 prompt_jsonl = sys.argv[1]160 in_glob = sys.argv[2] # Should probably be in a directory called "target/*.jpg".161 output_dir = sys.argv[3] # Should probably be a directory called "source".162 annotation_dir = ""163 if len(sys.argv) > 4:164 annotation_dir = sys.argv[4]165 main("generate_face_poses_checkpoint.json", prompt_jsonl, in_glob, output_dir, annotation_dir)166 else:167 print(f"""Usage:168 python {sys.argv[0]} prompt.jsonl target/*.jpg source/ [annotated/]169 source and target are slightly confusing in this context. We are writing the image names to prompt.jsonl, so 170 the naming system has to be consistent with what ControlNet expects. In ControlNet, the source is the input and171 target is the output. We are generating source images from targets in this application, so the second argument 172 should be a folder full of images. The third argument should be 'source', where the images should be places.173 Optionally, an 'annotated' directory can be provided. Augmented images will be placed here.174 175 A checkpoint file named 'generate_face_poses_checkpoint.json' will be created in the place where the script is 176 run. If a run is cancelled, it can be resumed from this checkpoint.177 178 If invoking the script from bash, do not forget to enclose globs with quotes. Example usage:179 `python ./tool_generate_face_poses.py ./face_prompt.jsonl "/home/josephcatrambone/training_data/data-mscoco/images/train2017/*" /home/josephcatrambone/training_data/data-mscoco/images/source_2017/`180 """)181 