diff --git a/README.md b/README.md index 329d546..31cce12 100644 --- a/README.md +++ b/README.md @@ -51,6 +51,8 @@ BMW with multi-alpha like [supermerger](https://github.com/hako-mikan/sd-webui-s |latent|RandomLatentImage|`INT`, `INT`, `INT`|`LATENT`|(width, height, batch_size)| |latent|VAEDecodeBatched|`LATENT`, `VAE`, `INT`|`IMAGE`|VAE decoding with specified batch size| |latent|VAEEncodeBatched|`IMAGE`, `VAE`, `INT`|`LATENT`|VAE encoding with specified batch size| +|latent|LatentToImage|`LATENT`, `FLOAT`, `FLOAT`|`IMAGE`|convert 4-ch latent tensor to 4 grayscale images| +|latent|LatentToHist|`LATENT`, `...`|`IMAGE`|create a histogram of the input latent| ### Sampling nodes diff --git a/__init__.py b/__init__.py index ccea139..a24f456 100644 --- a/__init__.py +++ b/__init__.py @@ -5,7 +5,8 @@ from .model.loader import StateDictLoader, Dict2Model from .model.iter import ModelIter, CLIPIter, VAEIter from .model.merge import StateDictMerger, StateDictMergerBlockWeighted from .model.merge2 import StateDictMergerBlockWeightedMulti -from .image import GridImage +from .image.image import GridImage +from .image.latenttoimage import LatentToImage, LatentToHist NODE_CLASS_MAPPINGS = { # latent @@ -16,6 +17,10 @@ NODE_CLASS_MAPPINGS = { 'VAEDecodeBatched': VAEDecodeBatched, 'VAEEncodeBatched': VAEEncodeBatched, + ## convert latent matrix to images + 'LatentToImage': LatentToImage, + 'LatentToHist': LatentToHist, + # sampling ## put parameters for sampler into a dict diff --git a/image.py b/image/image.py similarity index 100% rename from image.py rename to image/image.py diff --git a/image/latenttoimage.py b/image/latenttoimage.py new file mode 100644 index 0000000..7a8a51a --- /dev/null +++ b/image/latenttoimage.py @@ -0,0 +1,178 @@ +import colorsys +import torch +import torchvision.transforms.functional +import PIL.Image +import einops + +class LatentToImage: + + @classmethod + def INPUT_TYPES(cls): + return { + 'required': { + 'samples': ('LATENT',), + 'clamp_min': ('FLOAT', { 'default': -5.0, 'min': -100.0, 'max': 100.0, 'step': 0.01, }), + 'clamp_max': ('FLOAT', { 'default': 5.0, 'min': -100.0, 'max': 100.0, 'step': 0.01, }), + }, + #'optional': { + #} + } + + RETURN_TYPES = ('IMAGE',) + FUNCTION = 'execute' + + CATEGORY = 'latent' + + def execute( + self, + samples: dict, + clamp_min: float, + clamp_max: float, + ): + s: torch.Tensor = samples['samples'] + B, C, H, W = s.shape + assert C == 4 + + clamp_min = float(clamp_min) + clamp_max = float(clamp_max) + + if clamp_max < clamp_min: + clamp_min, clamp_max = clamp_max, clamp_min + + if abs(clamp_max - clamp_min) < 1e-3: + clamp_min = -5.0 + clamp_max = 5.0 + + s = s.clamp(min=clamp_min, max=clamp_max) + s = (s - clamp_min) / (clamp_max - clamp_min) + + images = [] + for b in range(B): + for c in range(C): + t = s[b,c,:,:] + #image = torchvision.transforms.functional.to_pil_image(t, mode='L') + #images.append(images) + rgb = torch.dstack([t,t,t]) + images.append(rgb.unsqueeze_(0)) + # (H,W) -> (B,H,W,C) + + return (torch.cat(images),) + +class LatentToHist: + + @classmethod + def INPUT_TYPES(cls): + return { + 'required': { + 'samples': ('LATENT',), + 'min_auto': (['Auto', 'Specified'],), + 'min_value': ('FLOAT', { 'default': -5.0, 'min': -100.0, 'max': 0.0, 'step': 0.01, }), + 'max_auto': (['Auto', 'Specified'],), + 'max_value': ('FLOAT', { 'default': 5.0, 'min': 0.0, 'max': 100.0, 'step': 0.01, }), + 'bin_auto': (['Auto', 'Specified'],), + 'bin_count': ('INT', { 'default': 10, 'min': 3, 'max': 1000, 'step': 1, }), + 'ymax_auto': (['Auto', 'Specified'],), + 'ymax': ('FLOAT', { 'default': 1.0, 'min': 0.01, 'max': 1.0, 'step': 0.01, }), + }, + } + + RETURN_TYPES = ('IMAGE',) + FUNCTION = 'execute' + + CATEGORY = 'latent' + + def execute( + self, + samples: dict, + min_auto: str, + min_value: float, + max_auto: str, + max_value: float, + bin_auto: str, + bin_count: int, + ymax_auto: str, + ymax: float, + ): + s: torch.Tensor = samples['samples'] + B, C, H, W = s.shape + assert C == 4 + + ss = s.view((B,C,-1)) + + def is_auto(v: str): + return v.lower() == 'auto' + + min_ = torch.min(ss).item() if is_auto(min_auto) else float(min_value) + max_ = torch.max(ss).item() if is_auto(max_auto) else float(max_value) + bins = 10 if is_auto(bin_auto) else int(bin_count) + + assert min_ < max_ + assert 3 <= bins + + hists = [ self.hist(ss[b,:,:], min_, max_, bins) for b in range(B) ] + + if is_auto(ymax_auto): + ymax = max([ torch.max(hist).item() for batch_hists in hists for hist, bin_edges in batch_hists ]) + ymax += 0.01 + else: + ymax = float(ymax) + + image_tensors = list(self.plot(hists, ymax)) + + images = torch.cat(image_tensors) + return (images,) + + def hist(self, batch_data: torch.Tensor, min_: float, max_: float, bins: int): + assert batch_data.dim() == 2 + + C, N = batch_data.shape + + assert C == 4 + + hists = [] + + for c in range(C): + hist, bin_edges = torch.histogram(batch_data[c,:], bins=bins, range=(min_, max_)) + hist.div_(N) + hists.append((hist, bin_edges)) + + return hists + + def plot(self, hists: list, ymax: float): + try: + from matplotlib import pyplot as plt + except Exception as e: + raise RuntimeError('LatentToHist requires matplotlib. Please install it.') + + def color(c: int, C: int): + h = c / C + return colorsys.hsv_to_rgb(h, 1.0, 1.0) + + for batch_hists in hists: + C = len(batch_hists) + + fig, ax = plt.subplots(1, 1, figsize=(6,6)) + plots = [] + for c, (hist, bin_edges) in enumerate(batch_hists): + plot = ax.plot( + bin_edges[:-1], hist, color=color(c,C), + linewidth=2, marker='o', markersize=4, + ) + plots.append(plot[0]) + + ax.legend(plots, [f'ch={c}' for c in range(C)], loc=2) + ax.set_ylim(0, ymax) + ax.grid(visible=True) + fig.canvas.draw() + + image = PIL.Image.frombytes('RGB', fig.canvas.get_width_height(), fig.canvas.tostring_rgb()) # type: ignore + + plt.close(fig) + + image = image.resize((512, 512)) + + image_tensor = torchvision.transforms.functional.to_tensor(image) + # (C, H, W) + + image_tensor = einops.rearrange(image_tensor, 'c h w -> h w c').unsqueeze(0) + yield image_tensor