diff --git a/.gitignore b/.gitignore new file mode 100644 index 0000000..82f9275 --- /dev/null +++ b/.gitignore @@ -0,0 +1,162 @@ +# Byte-compiled / optimized / DLL files +__pycache__/ +*.py[cod] +*$py.class + +# C extensions +*.so + +# Distribution / packaging +.Python +build/ +develop-eggs/ +dist/ +downloads/ +eggs/ +.eggs/ +lib/ +lib64/ +parts/ +sdist/ +var/ +wheels/ +share/python-wheels/ +*.egg-info/ +.installed.cfg +*.egg +MANIFEST + +# PyInstaller +# Usually these files are written by a python script from a template +# before PyInstaller builds the exe, so as to inject date/other infos into it. +*.manifest +*.spec + +# Installer logs +pip-log.txt +pip-delete-this-directory.txt + +# Unit test / coverage reports +htmlcov/ +.tox/ +.nox/ +.coverage +.coverage.* +.cache +nosetests.xml +coverage.xml +*.cover +*.py,cover +.hypothesis/ +.pytest_cache/ +cover/ + +# Translations +*.mo +*.pot + +# Django stuff: +*.log +local_settings.py +db.sqlite3 +db.sqlite3-journal + +# Flask stuff: +instance/ +.webassets-cache + +# Scrapy stuff: +.scrapy + +# Sphinx documentation +docs/_build/ + +# PyBuilder +.pybuilder/ +target/ + +# Jupyter Notebook +.ipynb_checkpoints + +# IPython +profile_default/ +ipython_config.py + +# pyenv +# For a library or package, you might want to ignore these files since the code is +# intended to run in multiple environments; otherwise, check them in: +# .python-version + +# pipenv +# According to pypa/pipenv#598, it is recommended to include Pipfile.lock in version control. +# However, in case of collaboration, if having platform-specific dependencies or dependencies +# having no cross-platform support, pipenv may install dependencies that don't work, or not +# install all needed dependencies. +#Pipfile.lock + +# poetry +# Similar to Pipfile.lock, it is generally recommended to include poetry.lock in version control. +# This is especially recommended for binary packages to ensure reproducibility, and is more +# commonly ignored for libraries. +# https://python-poetry.org/docs/basic-usage/#commit-your-poetrylock-file-to-version-control +#poetry.lock + +# pdm +# Similar to Pipfile.lock, it is generally recommended to include pdm.lock in version control. +#pdm.lock +# pdm stores project-wide configurations in .pdm.toml, but it is recommended to not include it +# in version control. +# https://pdm.fming.dev/latest/usage/project/#working-with-version-control +.pdm.toml +.pdm-python +.pdm-build/ + +# PEP 582; used by e.g. github.com/David-OConnor/pyflow and github.com/pdm-project/pdm +__pypackages__/ + +# Celery stuff +celerybeat-schedule +celerybeat.pid + +# SageMath parsed files +*.sage.py + +# Environments +.env +.venv +env/ +venv/ +ENV/ +env.bak/ +venv.bak/ + +# Spyder project settings +.spyderproject +.spyproject + +# Rope project settings +.ropeproject + +# mkdocs documentation +/site + +# mypy +.mypy_cache/ +.dmypy.json +dmypy.json + +# Pyre type checker +.pyre/ + +# pytype static type analyzer +.pytype/ + +# Cython debug symbols +cython_debug/ + +# PyCharm +# JetBrains specific template is maintained in a separate JetBrains.gitignore that can +# be found at https://github.com/github/gitignore/blob/main/Global/JetBrains.gitignore +# and can be added to the global gitignore or merged into this file. For a more nuclear +# option (not recommended) you can uncomment the following to ignore the entire idea folder. +#.idea/ diff --git a/__init__.py b/__init__.py new file mode 100644 index 0000000..7e59f2b --- /dev/null +++ b/__init__.py @@ -0,0 +1,23 @@ +"""TODO""" + +# pylint: disable=invalid-name + +from .advanced_tiling import ( + AdvancedTilingSettings, + AdvancedTiling, + AdvancedTilingVAEDecode, +) + +NODE_CLASS_MAPPINGS = { + "AdvancedTilingSettings": AdvancedTilingSettings, + "AdvancedTiling": AdvancedTiling, + "AdvancedTilingVAEDecode": AdvancedTilingVAEDecode, +} + +NODE_DISPLAY_NAME_MAPPINGS = { + "AdvancedTilingSettings": "Advanced Tiling Settings", + "AdvancedTiling": "Advanced Tiling", + "AdvancedTilingVAEDecode": "Advanced Tiling VAE Decode", +} + +__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"] diff --git a/advanced_tiling.py b/advanced_tiling.py new file mode 100644 index 0000000..018ee8d --- /dev/null +++ b/advanced_tiling.py @@ -0,0 +1,215 @@ +from typing import Optional +import functools +import copy + +from .modes import modes +from torch import Tensor +from torch.nn import Conv2d +from torch.nn import functional as F +from torch.nn.modules.utils import _pair +import numpy as np + + +class Settings: + """ + For representing tiling settings + """ + + def __init__(self, mode, rotation): + self.mode = mode + self.tiling_fn = modes[mode] + self.rotation = rotation + + def __hash__(self): + # We don't care about the tiling function, because it's determined by the mode + return hash((self.mode, self.rotation)) + + +@functools.cache +def calculate_mapping( + original_size: tuple[int, int], padded_size: tuple[int, int], settings: Settings +): + """ + Calculate mapping for pixels outside of the mask + + :param original_size: Original size of the image + :param padded_size: Padded size of the image + :param settings: Tiling settings + :return: Mapping of pixels + """ + + mapping = [] + for y in range(padded_size[1]): + for x in range(padded_size[0]): + (new_x, new_y) = settings.tiling_fn(x, y, original_size, padded_size) + mapping.append([x, y, new_x, new_y]) + return list(zip(*mapping)) + + +@functools.cache +def crop_image(image, settings: Settings): + """ + Crop image based on tiling settings + + :param image: Image to crop + :param settings: Tiling settings + :return: Cropped image + """ + + height, width = image.shape[1:3] + with_alpha = F.pad(image, (0, 1), "constant", 0) + print("size", height, width) + + for y in range(height): + for x in range(width): + # Calculate new coordinates + (new_x, new_y) = settings.tiling_fn(x, y, (width, height), (width, height)) + + # If coordinates match, it means we are in the mask + if new_x == x and new_y == y: + with_alpha[:, y, x, -1] = 1 + + return with_alpha + + +def patch_model(model, settings: Settings): + """ + TODO + """ + + # Patch all Conv2d layers + for layer in [layer for layer in model.modules() if isinstance(layer, Conv2d)]: + # pylint: disable=protected-access, no-value-for-parameter + layer._conv_forward = tiling_conv.__get__(layer, Conv2d) + layer.tiling_settings = settings + return model + + +def tiling_conv(self, input_tensor: Tensor, weight: Tensor, bias: Optional[Tensor]): + """ + TODO + """ + + # Pad input tensor + padded = F.pad( + input_tensor, + # pylint: disable=protected-access + self._reversed_padding_repeated_twice, + ) + # Calculate mapping + mapping = calculate_mapping( + (input_tensor.shape[-1], input_tensor.shape[-2]), + (padded.shape[-1], padded.shape[-2]), + self.tiling_settings, + ) + # Apply tiling + padded[:, :, mapping[1], mapping[0]] = padded[:, :, mapping[3], mapping[2]] + # Perform convolution + return F.conv2d( + padded, weight, bias, self.stride, _pair(0), self.dilation, self.groups + ) + + +class AdvancedTilingSettings: + """TODO""" + + # pylint: disable=invalid-name + + @classmethod + def INPUT_TYPES(cls): + """TODO""" + + return { + "required": { + "mode": (list(modes.keys()),), + "rotation": ( + "FLOAT", + {"default": 0.0, "min": 0.0, "max": 360.0, "step": 0.01}, + ), + }, + } + + RETURN_TYPES = ("ADVANCED_TILING_SETTINGS",) + RETURN_NAMES = ("SETTINGS",) + FUNCTION = "run" + + def run(self, mode, rotation): + """ + TODO + """ + + settings = Settings(mode, rotation) + + return (settings,) + + +class AdvancedTiling: + """ + Patches Conv2D layers in a model to perform tiling + """ + + # pylint: disable=invalid-name + + @classmethod + def INPUT_TYPES(cls): + """TODO""" + + return { + "required": { + "settings": ("ADVANCED_TILING_SETTINGS",), + "model": ("MODEL",), + }, + } + + CATEGORY = "conditioning" + RETURN_TYPES = ("MODEL",) + FUNCTION = "run" + + def run(self, settings, model): + """ + Does the actual patching of the model + """ + + model_copy = copy.deepcopy(model) + patch_model(model_copy.model, settings) + + return (model_copy,) + + +class AdvancedTilingVAEDecode: + """TODO""" + + # pylint: disable=invalid-name + + @classmethod + def INPUT_TYPES(cls): + """TODO""" + + return { + "required": { + "settings": ("ADVANCED_TILING_SETTINGS",), + "samples": ("LATENT",), + "vae": ("VAE",), + "crop": ("BOOLEAN", {"default": True}), + } + } + + RETURN_TYPES = ("IMAGE",) + FUNCTION = "run" + CATEGORY = "latent" + + def run(self, settings, samples, vae, crop): + """TODO""" + + print("settings", settings) + + vae_copy = copy.deepcopy(vae) + # Enable tiling + patch_model(vae_copy.first_stage_model, settings) + # Decode latents to image + image = vae_copy.decode(samples["samples"]) + if crop: + # Crop image based on tiling settings + image = crop_image(image, settings) + + return (image,) diff --git a/modes/__init__.py b/modes/__init__.py new file mode 100644 index 0000000..07e7df6 --- /dev/null +++ b/modes/__init__.py @@ -0,0 +1,13 @@ +""" +Collection of tiling modes +""" + +from .hex import hex_tiling +from .none import none_tiling + +modes = { + "None": none_tiling, + "Hexagon": hex_tiling, +} + +__all__ = ["modes"] diff --git a/modes/hex.py b/modes/hex.py new file mode 100644 index 0000000..8c4e9a7 --- /dev/null +++ b/modes/hex.py @@ -0,0 +1,92 @@ +""" +Hexagonal tiling implementation + +Some of this code is taken from excelent guide https://www.redblobgames.com/grids/hexagons/ +""" + +import math + + +def cube_to_axial(cube_coords: tuple[int, int, int]) -> tuple[int, int]: + """ + Convert cube coordinates to axial coordinates + + :param cube_coords: Cube coordinates + :return: Axial coordinates + """ + + return (cube_coords[0], cube_coords[1]) + + +def axial_to_cube(axial_coords: tuple[int, int]) -> tuple[int, int, int]: + """ + Convert axial coordinates to cube coordinates + + :param axial_coords: Axial coordinates + :return: Cube coordinates + """ + + q = axial_coords[0] + r = axial_coords[1] + s = -q - r + + return (q, r, s) + + +def axial_round(frac_coords: tuple[float, float]) -> tuple[int, int]: + """ + Round fractional axial coordinates to nearest axial coordinate + + :param frac_coords: Fractional axial coordinates + :return: Axial coordinates + """ + + return cube_to_axial(cube_round(axial_to_cube(frac_coords))) + + +def cube_round(frac_coords: tuple[float, float, float]) -> tuple[int, int, int]: + """ + Round fractional cube coordinates to nearest cube coordinate + + :param frac_coords: Fractional cube coordinates + :return: Cube coordinates + """ + + q = round(frac_coords[0]) + r = round(frac_coords[1]) + s = round(frac_coords[2]) + + q_diff = abs(q - frac_coords[0]) + r_diff = abs(r - frac_coords[1]) + s_diff = abs(s - frac_coords[2]) + + if q_diff > r_diff and q_diff > s_diff: + q = -r - s + elif r_diff > s_diff: + r = -q - s + else: + s = -q - r + + return (q, r, s) + + +def hex_tiling( + x: int, y: int, original_size: tuple[int, int], padded_size: tuple[int, int] +) -> tuple[int, int]: + ssize = padded_size[0] // 2 + size = original_size[0] // 2 + + q = (math.sqrt(3) / 3 * (x - ssize) - 1 / 3 * (y - ssize)) / size + r = (2 / 3 * (y - ssize)) / size + + rounded = axial_round((q, r)) + q -= rounded[0] + r -= rounded[1] + + xx = round(size * (math.sqrt(3) * q + (math.sqrt(3) / 2) * r)) + yy = round(size * ((3 / 2) * r)) + + xx = (xx + ssize) % padded_size[0] + yy = (yy + ssize) % padded_size[1] + + return (xx, yy) diff --git a/modes/none.py b/modes/none.py new file mode 100644 index 0000000..b08f395 --- /dev/null +++ b/modes/none.py @@ -0,0 +1,4 @@ +def none_tiling( + x: int, y: int, original_size: tuple[int, int], padded_size: tuple[int, int] +) -> tuple[int, int]: + return (x, y)