Team Ai
Apppublic

Dynamatrix/DiffBIR-OpenXLab

sourceHugging Faceapache-2.0updated 2y agoView on Hugging Face
0likes
util.py198 linesDownload Raw Back to ldm
1import importlib2 3import torch4from torch import optim5import numpy as np6 7from inspect import isfunction8from PIL import Image, ImageDraw, ImageFont9 10 11def log_txt_as_img(wh, xc, size=10):12    # wh a tuple of (width, height)13    # xc a list of captions to plot14    b = len(xc)15    txts = list()16    for bi in range(b):17        txt = Image.new("RGB", wh, color="white")18        draw = ImageDraw.Draw(txt)19        # font = ImageFont.truetype('font/DejaVuSans.ttf', size=size)20        font = ImageFont.load_default()21        nc = int(40 * (wh[0] / 256))22        lines = "\n".join(xc[bi][start:start + nc] for start in range(0, len(xc[bi]), nc))23 24        try:25            draw.text((0, 0), lines, fill="black", font=font)26        except UnicodeEncodeError:27            print("Cant encode string for logging. Skipping.")28 29        txt = np.array(txt).transpose(2, 0, 1) / 127.5 - 1.030        txts.append(txt)31    txts = np.stack(txts)32    txts = torch.tensor(txts)33    return txts34 35 36def ismap(x):37    if not isinstance(x, torch.Tensor):38        return False39    return (len(x.shape) == 4) and (x.shape[1] > 3)40 41 42def isimage(x):43    if not isinstance(x,torch.Tensor):44        return False45    return (len(x.shape) == 4) and (x.shape[1] == 3 or x.shape[1] == 1)46 47 48def exists(x):49    return x is not None50 51 52def default(val, d):53    if exists(val):54        return val55    return d() if isfunction(d) else d56 57 58def mean_flat(tensor):59    """60    https://github.com/openai/guided-diffusion/blob/27c20a8fab9cb472df5d6bdd6c8d11c8f430b924/guided_diffusion/nn.py#L8661    Take the mean over all non-batch dimensions.62    """63    return tensor.mean(dim=list(range(1, len(tensor.shape))))64 65 66def count_params(model, verbose=False):67    total_params = sum(p.numel() for p in model.parameters())68    if verbose:69        print(f"{model.__class__.__name__} has {total_params*1.e-6:.2f} M params.")70    return total_params71 72 73def instantiate_from_config(config):74    if not "target" in config:75        if config == '__is_first_stage__':76            return None77        elif config == "__is_unconditional__":78            return None79        raise KeyError("Expected key `target` to instantiate.")80    return get_obj_from_str(config["target"])(**config.get("params", dict()))81 82 83def get_obj_from_str(string, reload=False):84    module, cls = string.rsplit(".", 1)85    if reload:86        module_imp = importlib.import_module(module)87        importlib.reload(module_imp)88    return getattr(importlib.import_module(module, package=None), cls)89 90 91class AdamWwithEMAandWings(optim.Optimizer):92    # credit to https://gist.github.com/crowsonkb/65f7265353f403714fce3b2595e0b29893    def __init__(self, params, lr=1.e-3, betas=(0.9, 0.999), eps=1.e-8,  # TODO: check hyperparameters before using94                 weight_decay=1.e-2, amsgrad=False, ema_decay=0.9999,   # ema decay to match previous code95                 ema_power=1., param_names=()):96        """AdamW that saves EMA versions of the parameters."""97        if not 0.0 <= lr:98            raise ValueError("Invalid learning rate: {}".format(lr))99        if not 0.0 <= eps:100            raise ValueError("Invalid epsilon value: {}".format(eps))101        if not 0.0 <= betas[0] < 1.0:102            raise ValueError("Invalid beta parameter at index 0: {}".format(betas[0]))103        if not 0.0 <= betas[1] < 1.0:104            raise ValueError("Invalid beta parameter at index 1: {}".format(betas[1]))105        if not 0.0 <= weight_decay:106            raise ValueError("Invalid weight_decay value: {}".format(weight_decay))107        if not 0.0 <= ema_decay <= 1.0:108            raise ValueError("Invalid ema_decay value: {}".format(ema_decay))109        defaults = dict(lr=lr, betas=betas, eps=eps,110                        weight_decay=weight_decay, amsgrad=amsgrad, ema_decay=ema_decay,111                        ema_power=ema_power, param_names=param_names)112        super().__init__(params, defaults)113 114    def __setstate__(self, state):115        super().__setstate__(state)116        for group in self.param_groups:117            group.setdefault('amsgrad', False)118 119    @torch.no_grad()120    def step(self, closure=None):121        """Performs a single optimization step.122        Args:123            closure (callable, optional): A closure that reevaluates the model124                and returns the loss.125        """126        loss = None127        if closure is not None:128            with torch.enable_grad():129                loss = closure()130 131        for group in self.param_groups:132            params_with_grad = []133            grads = []134            exp_avgs = []135            exp_avg_sqs = []136            ema_params_with_grad = []137            state_sums = []138            max_exp_avg_sqs = []139            state_steps = []140            amsgrad = group['amsgrad']141            beta1, beta2 = group['betas']142            ema_decay = group['ema_decay']143            ema_power = group['ema_power']144 145            for p in group['params']:146                if p.grad is None:147                    continue148                params_with_grad.append(p)149                if p.grad.is_sparse:150                    raise RuntimeError('AdamW does not support sparse gradients')151                grads.append(p.grad)152 153                state = self.state[p]154 155                # State initialization156                if len(state) == 0:157                    state['step'] = 0158                    # Exponential moving average of gradient values159                    state['exp_avg'] = torch.zeros_like(p, memory_format=torch.preserve_format)160                    # Exponential moving average of squared gradient values161                    state['exp_avg_sq'] = torch.zeros_like(p, memory_format=torch.preserve_format)162                    if amsgrad:163                        # Maintains max of all exp. moving avg. of sq. grad. values164                        state['max_exp_avg_sq'] = torch.zeros_like(p, memory_format=torch.preserve_format)165                    # Exponential moving average of parameter values166                    state['param_exp_avg'] = p.detach().float().clone()167 168                exp_avgs.append(state['exp_avg'])169                exp_avg_sqs.append(state['exp_avg_sq'])170                ema_params_with_grad.append(state['param_exp_avg'])171 172                if amsgrad:173                    max_exp_avg_sqs.append(state['max_exp_avg_sq'])174 175                # update the steps for each param group update176                state['step'] += 1177                # record the step after step update178                state_steps.append(state['step'])179 180            optim._functional.adamw(params_with_grad,181                    grads,182                    exp_avgs,183                    exp_avg_sqs,184                    max_exp_avg_sqs,185                    state_steps,186                    amsgrad=amsgrad,187                    beta1=beta1,188                    beta2=beta2,189                    lr=group['lr'],190                    weight_decay=group['weight_decay'],191                    eps=group['eps'],192                    maximize=False)193 194            cur_ema_decay = min(ema_decay, 1 - state['step'] ** -ema_power)195            for param, ema_param in zip(params_with_grad, ema_params_with_grad):196                ema_param.mul_(cur_ema_decay).add_(param.float(), alpha=1 - cur_ema_decay)197 198        return loss