CoolFace
Apppublic

souging/TRELLIS_TextTo3D

sourceHugging Facemitupdated 1y agoView on Hugging Face
0likes
general_utils.py203 linesDownload Raw Back to utils
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