Team Ai
Modelpublic

CrucibleAI/ControlNetMediaPipeFace

sourceHugging Faceopenrailupdated 3y agoView on Hugging Face
576likes1.2kdownloads
tool_generate_face_poses.py181 linesDownload Raw Back to root
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