Fabrice-TIERCELIN commited on
Commit
5b2f74f
·
verified ·
1 Parent(s): 1f68d46

Delete utils.py

Browse files
Files changed (1) hide show
  1. utils.py +0 -172
utils.py DELETED
@@ -1,172 +0,0 @@
1
- import os
2
- import math
3
- import torch
4
- import logging
5
- import subprocess
6
- import numpy as np
7
- import torch.distributed as dist
8
-
9
- # from torch._six import inf
10
- from torch import inf
11
- from PIL import Image
12
- from typing import Union, Iterable
13
- from collections import OrderedDict
14
- from torch.utils.tensorboard import SummaryWriter
15
- from typing import Dict
16
- import torch_dct
17
-
18
- from diffusers.utils import is_bs4_available, is_ftfy_available
19
-
20
- import html
21
- import re
22
- import urllib.parse as ul
23
-
24
- if is_bs4_available():
25
- from bs4 import BeautifulSoup
26
-
27
- if is_ftfy_available():
28
- import ftfy
29
-
30
- import torch.fft as fft
31
-
32
- _tensor_or_tensors = Union[torch.Tensor, Iterable[torch.Tensor]]
33
-
34
-
35
- #################################################################################
36
- # Testing Utils #
37
- #################################################################################
38
-
39
- def find_model(model_name):
40
- """
41
- Finds a pre-trained model
42
- """
43
- assert os.path.isfile(model_name), f'Could not find DiT checkpoint at {model_name}'
44
- checkpoint = torch.load(model_name, map_location=lambda storage, loc: storage)
45
-
46
- if "ema" in checkpoint: # supports checkpoints from train.py
47
- print('Using ema ckpt!')
48
- checkpoint = checkpoint["ema"]
49
- else:
50
- checkpoint = checkpoint["model"]
51
- print("Using model ckpt!")
52
- return checkpoint
53
-
54
- def save_video_grid(video, nrow=None):
55
- b, t, h, w, c = video.shape
56
-
57
- if nrow is None:
58
- nrow = math.ceil(math.sqrt(b))
59
- ncol = math.ceil(b / nrow)
60
- padding = 1
61
- video_grid = torch.zeros((t, (padding + h) * nrow + padding,
62
- (padding + w) * ncol + padding, c), dtype=torch.uint8)
63
-
64
- # print(video_grid.shape)
65
- for i in range(b):
66
- r = i // ncol
67
- c = i % ncol
68
- start_r = (padding + h) * r
69
- start_c = (padding + w) * c
70
- video_grid[:, start_r:start_r + h, start_c:start_c + w] = video[i]
71
-
72
- return video_grid
73
-
74
- def save_videos_grid_tav(videos: torch.Tensor, path: str, rescale=False, nrow=None, fps=8):
75
- from einops import rearrange
76
- import imageio
77
- import torchvision
78
-
79
- b, _, _, _, _ = videos.shape
80
- if nrow is None:
81
- nrow = math.ceil(math.sqrt(b))
82
- videos = rearrange(videos, "b c t h w -> t b c h w")
83
- outputs = []
84
- for x in videos:
85
- x = torchvision.utils.make_grid(x, nrow=nrow)
86
- x = x.transpose(0, 1).transpose(1, 2).squeeze(-1)
87
- if rescale:
88
- x = (x + 1.0) / 2.0 # -1,1 -> 0,1
89
- x = (x * 255).numpy().astype(np.uint8)
90
- outputs.append(x)
91
-
92
- # os.makedirs(os.path.dirname(path), exist_ok=True)
93
- imageio.mimsave(path, outputs, fps=fps)
94
-
95
-
96
- #################################################################################
97
- # MMCV Utils #
98
- #################################################################################
99
-
100
-
101
- def collect_env():
102
- # Copyright (c) OpenMMLab. All rights reserved.
103
- from mmcv.utils import collect_env as collect_base_env
104
- from mmcv.utils import get_git_hash
105
- """Collect the information of the running environments."""
106
-
107
- env_info = collect_base_env()
108
- env_info['MMClassification'] = get_git_hash()[:7]
109
-
110
- for name, val in env_info.items():
111
- print(f'{name}: {val}')
112
-
113
- print(torch.cuda.get_arch_list())
114
- print(torch.version.cuda)
115
-
116
-
117
- #################################################################################
118
- # DCT Functions #
119
- #################################################################################
120
-
121
- def dct_low_pass_filter(dct_coefficients, percentage=0.3): # 2d [b c f h w]
122
- """
123
- Applies a low pass filter to the given DCT coefficients.
124
-
125
- :param dct_coefficients: 2D tensor of DCT coefficients
126
- :param percentage: percentage of coefficients to keep (between 0 and 1)
127
- :return: 2D tensor of DCT coefficients after applying the low pass filter
128
- """
129
- # Determine the cutoff indices for both dimensions
130
- cutoff_x = int(dct_coefficients.shape[-2] * percentage)
131
- cutoff_y = int(dct_coefficients.shape[-1] * percentage)
132
-
133
- # Create a mask with the same shape as the DCT coefficients
134
- mask = torch.zeros_like(dct_coefficients)
135
- # Set the top-left corner of the mask to 1 (the low-frequency area)
136
- mask[:, :, :, :cutoff_x, :cutoff_y] = 1
137
-
138
- return mask
139
-
140
- def normalize(tensor):
141
- """将Tensor归一化到[0, 1]范围内。"""
142
- min_val = tensor.min()
143
- max_val = tensor.max()
144
- normalized = (tensor - min_val) / (max_val - min_val)
145
- return normalized
146
-
147
- def denormalize(tensor, max_val_target, min_val_target):
148
- """将Tensor从[0, 1]范围反归一化到目标的[min_val_target, max_val_target]范围。"""
149
- denormalized = tensor * (max_val_target - min_val_target) + min_val_target
150
- return denormalized
151
-
152
- def exchanged_mixed_dct_freq(noise, base_content, LPF_3d, normalized=False):
153
- # noise dct
154
- noise_freq = torch_dct.dct_3d(noise, 'ortho')
155
-
156
- # frequency
157
- HPF_3d = 1 - LPF_3d
158
- noise_freq_high = noise_freq * HPF_3d
159
-
160
- # base frame dct
161
- base_content_freq = torch_dct.dct_3d(base_content, 'ortho')
162
-
163
- # base content low frequency
164
- base_content_freq_low = base_content_freq * LPF_3d
165
-
166
- # mixed frequency
167
- mixed_freq = base_content_freq_low + noise_freq_high
168
-
169
- # idct
170
- mixed_freq = torch_dct.idct_3d(mixed_freq, 'ortho')
171
-
172
- return mixed_freq