souging/TRELLIS_TextTo3D
0
1import re2import numpy as np3import cv24import torch5import contextlib6 7 8# Dictionary utils9def _dict_merge(dicta, dictb, prefix=''):10 """11 Merge two dictionaries.12 """13 assert isinstance(dicta, dict), 'input must be a dictionary'14 assert isinstance(dictb, dict), 'input must be a dictionary'15 dict_ = {}16 all_keys = set(dicta.keys()).union(set(dictb.keys()))17 for key in all_keys:18 if key in dicta.keys() and key in dictb.keys():19 if isinstance(dicta[key], dict) and isinstance(dictb[key], dict):20 dict_[key] = _dict_merge(dicta[key], dictb[key], prefix=f'{prefix}.{key}')21 else:22 raise ValueError(f'Duplicate key {prefix}.{key} found in both dictionaries. Types: {type(dicta[key])}, {type(dictb[key])}')23 elif key in dicta.keys():24 dict_[key] = dicta[key]25 else:26 dict_[key] = dictb[key]27 return dict_28 29 30def dict_merge(dicta, dictb):31 """32 Merge two dictionaries.33 """34 return _dict_merge(dicta, dictb, prefix='')35 36 37def dict_foreach(dic, func, special_func={}):38 """39 Recursively apply a function to all non-dictionary leaf values in a dictionary.40 """41 assert isinstance(dic, dict), 'input must be a dictionary'42 for key in dic.keys():43 if isinstance(dic[key], dict):44 dic[key] = dict_foreach(dic[key], func)45 else:46 if key in special_func.keys():47 dic[key] = special_func[key](dic[key])48 else:49 dic[key] = func(dic[key])50 return dic51 52 53def dict_reduce(dicts, func, special_func={}):54 """55 Reduce a list of dictionaries. Leaf values must be scalars.56 """57 assert isinstance(dicts, list), 'input must be a list of dictionaries'58 assert all([isinstance(d, dict) for d in dicts]), 'input must be a list of dictionaries'59 assert len(dicts) > 0, 'input must be a non-empty list of dictionaries'60 all_keys = set([key for dict_ in dicts for key in dict_.keys()])61 reduced_dict = {}62 for key in all_keys:63 vlist = [dict_[key] for dict_ in dicts if key in dict_.keys()]64 if isinstance(vlist[0], dict):65 reduced_dict[key] = dict_reduce(vlist, func, special_func)66 else:67 if key in special_func.keys():68 reduced_dict[key] = special_func[key](vlist)69 else:70 reduced_dict[key] = func(vlist)71 return reduced_dict72 73 74def dict_any(dic, func):75 """76 Recursively apply a function to all non-dictionary leaf values in a dictionary.77 """78 assert isinstance(dic, dict), 'input must be a dictionary'79 for key in dic.keys():80 if isinstance(dic[key], dict):81 if dict_any(dic[key], func):82 return True83 else:84 if func(dic[key]):85 return True86 return False87 88 89def dict_all(dic, func):90 """91 Recursively apply a function to all non-dictionary leaf values in a dictionary.92 """93 assert isinstance(dic, dict), 'input must be a dictionary'94 for key in dic.keys():95 if isinstance(dic[key], dict):96 if not dict_all(dic[key], func):97 return False98 else:99 if not func(dic[key]):100 return False101 return True102 103 104def dict_flatten(dic, sep='.'):105 """106 Flatten a nested dictionary into a dictionary with no nested dictionaries.107 """108 assert isinstance(dic, dict), 'input must be a dictionary'109 flat_dict = {}110 for key in dic.keys():111 if isinstance(dic[key], dict):112 sub_dict = dict_flatten(dic[key], sep=sep)113 for sub_key in sub_dict.keys():114 flat_dict[str(key) + sep + str(sub_key)] = sub_dict[sub_key]115 else:116 flat_dict[key] = dic[key]117 return flat_dict118 119 120# Context utils121@contextlib.contextmanager122def nested_contexts(*contexts):123 with contextlib.ExitStack() as stack:124 for ctx in contexts:125 stack.enter_context(ctx())126 yield127 128 129# Image utils130def make_grid(images, nrow=None, ncol=None, aspect_ratio=None):131 num_images = len(images)132 if nrow is None and ncol is None:133 if aspect_ratio is not None:134 nrow = int(np.round(np.sqrt(num_images / aspect_ratio)))135 else:136 nrow = int(np.sqrt(num_images))137 ncol = (num_images + nrow - 1) // nrow138 elif nrow is None and ncol is not None:139 nrow = (num_images + ncol - 1) // ncol140 elif nrow is not None and ncol is None:141 ncol = (num_images + nrow - 1) // nrow142 else:143 assert nrow * ncol >= num_images, 'nrow * ncol must be greater than or equal to the number of images'144 145 if images[0].ndim == 2:146 grid = np.zeros((nrow * images[0].shape[0], ncol * images[0].shape[1]), dtype=images[0].dtype)147 else:148 grid = np.zeros((nrow * images[0].shape[0], ncol * images[0].shape[1], images[0].shape[2]), dtype=images[0].dtype)149 for i, img in enumerate(images):150 row = i // ncol151 col = i % ncol152 grid[row * img.shape[0]:(row + 1) * img.shape[0], col * img.shape[1]:(col + 1) * img.shape[1]] = img153 return grid154 155 156def notes_on_image(img, notes=None):157 img = np.pad(img, ((0, 32), (0, 0), (0, 0)), 'constant', constant_values=0)158 img = cv2.cvtColor(img, cv2.COLOR_RGB2BGR)159 if notes is not None:160 img = cv2.putText(img, notes, (0, img.shape[0] - 4), cv2.FONT_HERSHEY_SIMPLEX, 1, (255, 255, 255), 1)161 img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)162 return img163 164 165def save_image_with_notes(img, path, notes=None):166 """167 Save an image with notes.168 """169 if isinstance(img, torch.Tensor):170 img = img.cpu().numpy().transpose(1, 2, 0)171 if img.dtype == np.float32 or img.dtype == np.float64:172 img = np.clip(img * 255, 0, 255).astype(np.uint8)173 img = notes_on_image(img, notes)174 cv2.imwrite(path, cv2.cvtColor(img, cv2.COLOR_RGB2BGR))175 176 177# debug utils178 179def atol(x, y):180 """181 Absolute tolerance.182 """183 return torch.abs(x - y)184 185 186def rtol(x, y):187 """188 Relative tolerance.189 """190 return torch.abs(x - y) / torch.clamp_min(torch.maximum(torch.abs(x), torch.abs(y)), 1e-12)191 192 193# print utils194def indent(s, n=4):195 """196 Indent a string.197 """198 lines = s.split('\n')199 for i in range(1, len(lines)):200 lines[i] = ' ' * n + lines[i]201 return '\n'.join(lines)202 203 