mirror of
https://github.com/vladmandic/automatic
synced 2026-09-20 09:38:23 +02:00
@@ -74,21 +74,21 @@ def edge_detect_for_pixelart(image: PipelineImageInput, image_weight: float = 1.
|
||||
return new_image
|
||||
|
||||
|
||||
def get_dct_harmonics(N: int, device: torch.device) -> torch.FloatTensor:
|
||||
def get_dct_harmonics(N: int, device: torch.device) -> torch.Tensor:
|
||||
k = torch.arange(N, dtype=torch.float32, device=device)
|
||||
spatial = torch.add(1, k.unsqueeze(1), alpha=2)
|
||||
spectral = k.unsqueeze(0) * (torch.pi / (2 * N))
|
||||
return torch.cos(torch.mm(spatial, spectral))
|
||||
|
||||
|
||||
def get_dct_norm(N: int, device: torch.device) -> torch.FloatTensor:
|
||||
def get_dct_norm(N: int, device: torch.device) -> torch.Tensor:
|
||||
n = torch.ones((N, 1), dtype=torch.float32, device=device)
|
||||
n[0, 0] = 1 / math.sqrt(2)
|
||||
n = torch.mm(n, n.t())
|
||||
return n
|
||||
|
||||
|
||||
def dct_2d(x: torch.FloatTensor, norm: str="ortho") -> torch.FloatTensor:
|
||||
def dct_2d(x: torch.FloatTensor, norm: str="ortho") -> torch.Tensor:
|
||||
x_shape = x.shape
|
||||
N = x_shape[-1]
|
||||
x = x.contiguous().view(-1, N, N)
|
||||
@@ -102,7 +102,7 @@ def dct_2d(x: torch.FloatTensor, norm: str="ortho") -> torch.FloatTensor:
|
||||
return coeff
|
||||
|
||||
|
||||
def idct_2d(coeff: torch.FloatTensor, norm: str="ortho") -> torch.FloatTensor:
|
||||
def idct_2d(coeff: torch.FloatTensor, norm: str="ortho") -> torch.Tensor:
|
||||
x_shape = coeff.shape
|
||||
N = x_shape[-1]
|
||||
coeff = coeff.contiguous().view(-1, N, N)
|
||||
@@ -116,7 +116,7 @@ def idct_2d(coeff: torch.FloatTensor, norm: str="ortho") -> torch.FloatTensor:
|
||||
return x
|
||||
|
||||
|
||||
def encode_single_channel_dct_2d(img: torch.FloatTensor, block_size: int=16, norm: str="ortho") -> torch.FloatTensor:
|
||||
def encode_single_channel_dct_2d(img: torch.FloatTensor, block_size: int=16, norm: str="ortho") -> torch.Tensor:
|
||||
batch_size, height, width = img.shape
|
||||
h_blocks = int(height//block_size)
|
||||
w_blocks = int(width//block_size)
|
||||
@@ -130,7 +130,7 @@ def encode_single_channel_dct_2d(img: torch.FloatTensor, block_size: int=16, nor
|
||||
return dct_tensor
|
||||
|
||||
|
||||
def decode_single_channel_dct_2d(img: torch.FloatTensor, norm: str="ortho") -> torch.FloatTensor:
|
||||
def decode_single_channel_dct_2d(img: torch.FloatTensor, norm: str="ortho") -> torch.Tensor:
|
||||
batch_size, combined_block_size, h_blocks, w_blocks = img.shape
|
||||
block_size = int(math.sqrt(combined_block_size))
|
||||
height = int(h_blocks*block_size)
|
||||
@@ -142,19 +142,19 @@ def decode_single_channel_dct_2d(img: torch.FloatTensor, norm: str="ortho") -> t
|
||||
return img_tensor
|
||||
|
||||
|
||||
def rgb_to_ycbcr_tensor(image: torch.ByteTensor) -> torch.FloatTensor:
|
||||
def rgb_to_ycbcr_tensor(image: torch.ByteTensor) -> torch.Tensor:
|
||||
rgb_weights = torch.tensor([[0.002345098, -0.001323419, 0.003921569], [0.004603922, -0.00259815, -0.003283824], [0.000894118, 0.003921569, -0.000637744]], device=image.device)
|
||||
ycbcr = torch.einsum("cv,...chw->...vhw", [rgb_weights, image.permute(0,3,1,2).to(dtype=torch.float32)])
|
||||
ycbcr[:,0,:,:] = ycbcr[:,0,:,:].add(-1)
|
||||
return ycbcr
|
||||
|
||||
|
||||
def ycbcr_tensor_to_rgb(ycbcr: torch.FloatTensor) -> torch.ByteTensor:
|
||||
def ycbcr_tensor_to_rgb(ycbcr: torch.FloatTensor) -> torch.Tensor:
|
||||
ycbcr_weights = torch.tensor([[127.5, 127.5, 127.5], [0, -43.877376465, 225.93], [178.755, -91.052376465, 0]], device=ycbcr.device)
|
||||
return torch.einsum("cv,...chw->...vhw", [ycbcr_weights, ycbcr]).add(127.5).round().clamp(0,255).permute(0,2,3,1).to(dtype=torch.uint8)
|
||||
|
||||
|
||||
def encode_jpeg_tensor(img: torch.FloatTensor, block_size: int=16, cbcr_downscale: int=2, norm: str="ortho") -> torch.FloatTensor:
|
||||
def encode_jpeg_tensor(img: torch.FloatTensor, block_size: int=16, cbcr_downscale: int=2, norm: str="ortho") -> torch.Tensor:
|
||||
img = img[:, :, :(img.shape[-2]//block_size)*block_size, :(img.shape[-1]//block_size)*block_size] # crop to a multiply of block_size
|
||||
cbcr_block_size = block_size//cbcr_downscale
|
||||
_, _, height, width = img.shape
|
||||
@@ -165,7 +165,7 @@ def encode_jpeg_tensor(img: torch.FloatTensor, block_size: int=16, cbcr_downscal
|
||||
return torch.cat([y,cb,cr], dim=1)
|
||||
|
||||
|
||||
def decode_jpeg_tensor(jpeg_img: torch.FloatTensor, block_size: int=16, cbcr_downscale: int=2, norm: str="ortho") -> torch.FloatTensor:
|
||||
def decode_jpeg_tensor(jpeg_img: torch.FloatTensor, block_size: int=16, cbcr_downscale: int=2, norm: str="ortho") -> torch.Tensor:
|
||||
_, _, h_blocks, w_blocks = jpeg_img.shape
|
||||
y_block_size = block_size*block_size
|
||||
cbcr_block_size = int((block_size//cbcr_downscale) ** 2)
|
||||
@@ -181,7 +181,7 @@ def decode_jpeg_tensor(jpeg_img: torch.FloatTensor, block_size: int=16, cbcr_dow
|
||||
return torch.stack([y,cb,cr], dim=1)
|
||||
|
||||
|
||||
def process_image_input(images: PipelineImageInput) -> torch.ByteTensor:
|
||||
def process_image_input(images: PipelineImageInput) -> torch.Tensor:
|
||||
if isinstance(images, list):
|
||||
combined_images = []
|
||||
for img in images:
|
||||
@@ -235,7 +235,7 @@ class JPEGEncoder(ImageProcessingMixin, ConfigMixin):
|
||||
self.latents_mean = latents_mean
|
||||
super().__init__()
|
||||
|
||||
def encode(self, images: PipelineImageInput, device: str="cpu") -> torch.FloatTensor:
|
||||
def encode(self, images: PipelineImageInput, device: str="cpu") -> torch.Tensor:
|
||||
"""
|
||||
Encode RGB 0-255 image to JPEG Latents.
|
||||
|
||||
|
||||
Reference in New Issue
Block a user