File size: 5,365 Bytes
635195b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
import os
import cv2
import random
import json
import torch
from PIL import Image
from realesrgan import RealESRGANer
from basicsr.archs.rrdbnet_arch import RRDBNet
from unet import DiffusionModel


def get_prompts(categories: list, use_training_prompts: bool):
    rand_category = random.choice(categories).lower()

    if use_training_prompts:
        with open("training_prompts.json", "r") as f:
            training_prompts_dict = json.load(f)

        return random.choice(training_prompts_dict[rand_category])


    with open("test_prompts.json", "r") as f:
        new_prompts_dict = json.load(f)

    return random.choice(new_prompts_dict[rand_category])



def get_vocab(consecutive_words: int) -> list[str]:
    mapping = {
        1: "consecutive_words/single_word_f10.txt",
        2: "consecutive_words/double_word_f10.txt",
        3: "consecutive_words/triple_word_f10.txt",
        4: "consecutive_words/quad_word_f10.txt",
        5: "consecutive_words/penta_word_f10.txt",
    }

    with open(mapping[consecutive_words], "r") as f:
        vocab = f.read().split("\n")
        vocab = [" ".join(w.split(" ")[:-1]) for w in vocab if len(w) > 0]

    return vocab


def load_esrgan(esrgan_type: str, device: str):
    mapping = {
        "2x": "real_esrgan_weights/RealESRGAN_x2plus.pth",
        "4x": "real_esrgan_weights/RealESRGAN_x4plus.pth",
    }
    esr_scale = 2 if esrgan_type == "2x" else 4

    RRDBNet_model = RRDBNet(num_in_ch=3, num_out_ch=3, num_feat=64, num_block=23, num_grow_ch=32, scale=esr_scale)
    upsampler = RealESRGANer(
        scale=esr_scale,
        model_path=mapping[esrgan_type],
        model=RRDBNet_model,
        half=True if device == "cuda" else False  # Use FP16 on CUDA
    )

    return upsampler


def load_diffusion_model(initial_load: bool, diffusion_model_path: str, device: str):
    if initial_load:  # Meaning this is the initial loading of model (Default load)
        weights = [p for p in os.listdir("diffusion_model_weights") if p.endswith(".pth")]
        try:
            # Checks if model files still follows the expected name format of 'epoch_<ep#>_model_tl_<tl#>_vl_<vl#>.pth'
            # For example, 'epoch_150_model_tl_165_vl_183.pth'
            weights = sorted(weights, key=lambda x: int(x.split("_")[1]))
            diffusion_model_path = weights[-1]  # Use the final ("best") one
        except Exception as e:
            # Incase users changed the filename, just randomly select one as default, until user manually selects another one
            diffusion_model_path = random.choice(weights)
            print("Error in model weights selection due to incorrect filename formatting. Selecting random model"
                  f"\n{diffusion_model_path=}")


    diffusion_model_path = os.path.join("diffusion_model_weights", diffusion_model_path)

    diffusion_model_dict = torch.load(diffusion_model_path, weights_only=False, map_location=device)
    unet_params = diffusion_model_dict["unet_params"]
    diffusion_params = diffusion_model_dict["diffusion_params"]

    diffusion_model = DiffusionModel(unet_params=unet_params, diffusion_params=diffusion_params)
    diffusion_model.load_state_dict(diffusion_model_dict["model_state_dict"])
    diffusion_model.to(device)

    return diffusion_model, diffusion_model_path


def save_images(img_tensors) -> None:
    assert len(img_tensors.shape) == 4, f"Expected img_tensors to be of shape (b, w, h, c) instead got {img_tensors.shape=}"

    downloads_path = "./flask_outputs/batch_gen"

    for it in img_tensors:
        # Iterating over tensor like this would result in (w, h, c)
        img = Image.fromarray(it.to(torch.uint8).numpy())
        filename = get_unique_filename(downloads_path, "custom_img.png")
        img.save(os.path.join(downloads_path, filename))


def create_denoising_video(frame_rate):
    # Set parameters
    image_dir = "./flask_outputs/denoising_temp"
    output_dir = "./flask_outputs/denoising_outputs"
    output_path = get_unique_filename(output_dir, "output.mp4")
    output_path = os.path.join(output_dir, output_path)

    # Get list of images sorted by filename
    images = sorted([img for img in os.listdir(image_dir)], key=lambda x: int(x.split(".")[0]), reverse=True)
    images = [os.path.join(image_dir, img) for img in images]

    # Read the first image to get dimensions
    frame = cv2.imread(images[0])
    height, width, layers = frame.shape

    # Define video writer
    fourcc = cv2.VideoWriter_fourcc(*"mp4v")  # Codec
    video = cv2.VideoWriter(output_path, fourcc, frame_rate, (width, height))

    # Add images to video
    for img_path in images:
        frame = cv2.imread(img_path)
        video.write(frame)

    # Release video writer
    video.release()
    cv2.destroyAllWindows()


def get_unique_filename(directory, filename):
    # Prevent overwriting previous files. Appends like (x) suffix for filenames as commonly seen
    base, ext = os.path.splitext(filename)
    counter = 1
    new_filename = filename

    while os.path.exists(os.path.join(directory, new_filename)):
        new_filename = f"{base} ({counter}){ext}"
        counter += 1

    return new_filename


if __name__ == "__main__":
    ...