source code and requirements.
Should be working but might contain some old comments or unused code. Will try to add documentation and sample workflows next week.
This commit is contained in:
+25
@@ -0,0 +1,25 @@
|
||||
from .nodes import *
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"WFC_SampleNode_BMad": WFC_SampleNode,
|
||||
"WFC_Generate_BMad": WFC_GenerateNode,
|
||||
"WFC_Encode_BMad": WFC_Encode,
|
||||
"WFC_Decode_BMad": WFC_Decode,
|
||||
"WFC_CustomTemperature_Bmad": WFC_CustomTemperature,
|
||||
"WFC_CustomValueWeights_Bmad": WFC_CustomValueWeights,
|
||||
"WFC_EmptyState_Bmad": WFC_EmptyState,
|
||||
"WFC_Filter_Bmad": WFC_Filter,
|
||||
"WFC_GenParallel_Bmad": WFC_GenParallel,
|
||||
}
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"WFC_SampleNode_BMad": "Sample (WFC)",
|
||||
"WFC_Generate_BMad": "Generate (WFC)",
|
||||
"WFC_Encode_BMad": "Encode (WFC)",
|
||||
"WFC_Decode_BMad": "Decode (WFC)",
|
||||
"WFC_CustomTemperature_Bmad": "Custom Temperature Config (WFC)",
|
||||
"WFC_CustomValueWeights_Bmad": "Custom Value Weights (WFC)",
|
||||
"WFC_EmptyState_Bmad": "Empty State (WFC)",
|
||||
"WFC_Filter_Bmad": "Filter (WFC)",
|
||||
"WFC_GenParallel_Bmad": "Parallel Multi Gen. (WFC)",
|
||||
}
|
||||
@@ -0,0 +1,340 @@
|
||||
from _operator import xor
|
||||
from py_search.informed import best_first_search
|
||||
from comfy import utils
|
||||
from comfy.model_management import throw_exception_if_processing_interrupted, processing_interrupted
|
||||
from .wcf import *
|
||||
|
||||
|
||||
def waiting_loop(abort_loop_event, interruption_proxy: ValueProxy, pbar: utils.ProgressBar, ticker_proxy: ValueProxy, total_steps):
|
||||
"""
|
||||
Listens for interrupts and propagates to Problem running using interruption_proxy.
|
||||
Updates progress_bar via ticker_proxy, updated within Problem instances.
|
||||
|
||||
@param abort_loop_event: to be triggered in the main thread once the problem(s) have been solved
|
||||
@param interruption_proxy: proxy value used to cancel the search(es)
|
||||
@param pbar: comfyui progress bar to update every 100 milliseconds
|
||||
@param ticker_proxy: the max depth so far (or sum of max depths in case of many problems)
|
||||
@param total_steps: the total number of nodes to process in the problem(s)
|
||||
"""
|
||||
from time import sleep
|
||||
while not abort_loop_event.is_set():
|
||||
sleep(.1) # pause for 1 second
|
||||
if processing_interrupted():
|
||||
interruption_proxy.set(True)
|
||||
return
|
||||
pbar.update_absolute(ticker_proxy.get(), total_steps)
|
||||
|
||||
|
||||
class WFC_SampleNode:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required":
|
||||
{
|
||||
"img_batch": ("IMAGE",),
|
||||
"tile_width": ("INT", {"default": 32, "min": 1, "max": 128}),
|
||||
"tile_height": ("INT", {"default": 32, "min": 1, "max": 128}),
|
||||
"output_tiles": ("BOOLEAN", {"default": False})
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("WFC_Sample", "IMAGE",)
|
||||
RETURN_NAMES = ("sample", "unique_tiles",)
|
||||
FUNCTION = "compute"
|
||||
CATEGORY = "Bmad/WFC"
|
||||
|
||||
def compute(self, img_batch, tile_width, tile_height, output_tiles):
|
||||
import torch
|
||||
|
||||
samples = [np.clip(255. * img_batch[i].cpu().numpy().squeeze(), 0, 255).astype(np.uint8)
|
||||
for i in range(img_batch.shape[0])]
|
||||
sample = WFC_Sample(samples, tile_width, tile_height)
|
||||
|
||||
if output_tiles:
|
||||
tiles = [torch.from_numpy(tile.astype(np.float32) / 255.0).unsqueeze(0)
|
||||
for tile, freq in sample.get_tile_data().values()]
|
||||
tiles = torch.concat(tiles)
|
||||
else:
|
||||
tiles = torch.empty((1, 1, 1))
|
||||
|
||||
return (sample, tiles,)
|
||||
|
||||
|
||||
class WFC_GenerateNode:
|
||||
@staticmethod
|
||||
def NODE_INPUT_TYPES():
|
||||
return {
|
||||
"required":
|
||||
{
|
||||
"sample": ("WFC_Sample",),
|
||||
"starting_state": ("WFC_State",),
|
||||
"seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
|
||||
"max_freq_adjust": ("FLOAT", {"default": 1, "min": 0, "max": 1, "step": .01}),
|
||||
"use_8_cardinals": ("BOOLEAN", {"default": False}),
|
||||
"plateau_check_interval": ("INT", {"default": -1, "min": -1, "max": 10000}),
|
||||
},
|
||||
"optional":
|
||||
{
|
||||
"custom_temperature_config": ("WFC_TemperatureConfig",),
|
||||
"custom_node_value_config": ("WFC_NodeValueConfig",)
|
||||
}
|
||||
}
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return s.NODE_INPUT_TYPES()
|
||||
|
||||
RETURN_TYPES = ("WFC_State",)
|
||||
RETURN_NAMES = ("state", "unique_tiles",)
|
||||
FUNCTION = "compute"
|
||||
CATEGORY = "Bmad/WFC"
|
||||
|
||||
def compute(self, custom_temperature_config=None, custom_node_value_config=None, **kwargs):
|
||||
from multiprocessing import Manager
|
||||
from multiprocessing.managers import ValueProxy
|
||||
from threading import Event, Thread
|
||||
|
||||
if custom_temperature_config is not None:
|
||||
kwargs.update(custom_temperature_config)
|
||||
|
||||
if custom_node_value_config is not None:
|
||||
kwargs.update(custom_node_value_config)
|
||||
|
||||
# prepare stuff to process interrupts & update bar
|
||||
# TODO count is also done inside Problem, maybe should use as optional arg to avoid repeating the operation
|
||||
ss = kwargs["starting_state"]
|
||||
total_tiles_to_proc = ss.size-np.count_nonzero(ss)
|
||||
manager = Manager()
|
||||
stop: ValueProxy = manager.Value('b', False) # Set to True to abort
|
||||
ticker: ValueProxy = manager.Value('i', 0) # counts max depth increments
|
||||
kwargs.update({"stop_proxy": stop})
|
||||
kwargs.update({"ticker_proxy": ticker})
|
||||
finished_event = Event()
|
||||
pbar: utils.ProgressBar = utils.ProgressBar(total_tiles_to_proc)
|
||||
|
||||
t = Thread(target=waiting_loop, args=(finished_event, stop, pbar, ticker, total_tiles_to_proc))
|
||||
t.start()
|
||||
|
||||
problem = WFC_Problem(**kwargs)
|
||||
try:
|
||||
next(best_first_search(problem, graph=True)) # find 1st solution
|
||||
except InterruptedError:
|
||||
pass
|
||||
except StopIteration:
|
||||
print("Exhausted all possibilities without finding a complete solution ; or some irregularity occurred.")
|
||||
finally:
|
||||
finished_event.set()
|
||||
if stop.get():
|
||||
throw_exception_if_processing_interrupted()
|
||||
|
||||
result = problem.get_solution_state()
|
||||
return (result,)
|
||||
|
||||
|
||||
class WFC_Encode:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required":
|
||||
{
|
||||
"img": ("IMAGE",),
|
||||
"sample": ("WFC_Sample",),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("WFC_State",)
|
||||
RETURN_NAMES = ("state",)
|
||||
FUNCTION = "compute"
|
||||
CATEGORY = "Bmad/WFC"
|
||||
|
||||
def compute(self, img, sample: WFC_Sample):
|
||||
samples = [np.clip(255. * img[i].cpu().numpy().squeeze(), 0, 255).astype(np.uint8) for i in range(img.shape[0])]
|
||||
encoded = sample.img_to_tile_encoded_world(samples[0]) # no batch enconding, only a single image is encoded
|
||||
return (encoded,)
|
||||
|
||||
|
||||
class WFC_Decode:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required":
|
||||
{
|
||||
"state": ("WFC_State",),
|
||||
"sample": ("WFC_Sample",),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE", "MASK")
|
||||
FUNCTION = "compute"
|
||||
CATEGORY = "Bmad/WFC"
|
||||
|
||||
def compute(self, state, sample: WFC_Sample):
|
||||
import torch
|
||||
|
||||
img, mask = sample.tile_encoded_to_img(state)
|
||||
img = torch.from_numpy(img.astype(np.float32) / 255.0).unsqueeze(0)
|
||||
mask = torch.from_numpy(mask.astype(np.float32) / 255.0).unsqueeze(0)
|
||||
return (img, mask,)
|
||||
|
||||
|
||||
class WFC_CustomTemperature:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required":
|
||||
{
|
||||
"starting_temperature": ("INT", {"default": 50, "min": 0, "max": 99}),
|
||||
"min_min_temperature": ("INT", {"default": 0, "min": 0, "max": 99}),
|
||||
"max_min_temperature": ("INT", {"default": 80, "min": 0, "max": 99}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("WFC_TemperatureConfig",)
|
||||
RETURN_NAMES = ("temperature",)
|
||||
FUNCTION = "send"
|
||||
CATEGORY = "Bmad/WFC"
|
||||
|
||||
def send(self, **kwargs):
|
||||
return (kwargs,)
|
||||
|
||||
|
||||
class WFC_CustomValueWeights:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required":
|
||||
{
|
||||
"reverse_depth_w": ("FLOAT", {"default": 1, "min": 0, "max": 10, "step": .001}),
|
||||
"node_cost_w": ("FLOAT", {"default": 1, "min": 0, "max": 10, "step": .001}),
|
||||
"path_entropy_average_w": ("FLOAT", {"default": 0, "min": 0, "max": 10, "step": .001}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("WFC_NodeValueConfig",)
|
||||
RETURN_NAMES = ("weights",)
|
||||
FUNCTION = "send"
|
||||
CATEGORY = "Bmad/WFC"
|
||||
|
||||
def send(self, **kwargs):
|
||||
return (kwargs,)
|
||||
|
||||
|
||||
class WFC_EmptyState:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required":
|
||||
{
|
||||
"width": ("INT", {"default": 16, "min": 4, "max": 128}),
|
||||
"height": ("INT", {"default": 16, "min": 4, "max": 128}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("WFC_State",)
|
||||
RETURN_NAMES = ("state",)
|
||||
FUNCTION = "create"
|
||||
CATEGORY = "Bmad/WFC"
|
||||
|
||||
def create(self, width, height):
|
||||
return (np.zeros((width, height)),)
|
||||
|
||||
|
||||
class WFC_Filter:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required":
|
||||
{
|
||||
"state": ("WFC_State",),
|
||||
"tiles_batch": ("IMAGE",),
|
||||
"invert": ("BOOLEAN", {"default": False}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("WFC_State",)
|
||||
RETURN_NAMES = ("state",)
|
||||
FUNCTION = "create"
|
||||
CATEGORY = "Bmad/WFC"
|
||||
|
||||
def create(self, state: ndarray, tiles_batch, invert):
|
||||
to_filter = [WFC_Sample.tile_to_hash(
|
||||
np.clip(255. * tiles_batch[i].cpu().numpy().squeeze(), 0, 255).astype(np.uint8))
|
||||
for i in range(tiles_batch.shape[0])]
|
||||
new_state = [t if xor(t in to_filter, invert) else 0 for t in state.flatten()]
|
||||
new_state = np.array(new_state).reshape(state.shape)
|
||||
return (new_state,)
|
||||
|
||||
|
||||
def generate_single(i_kwargs):
|
||||
problem = WFC_Problem(**i_kwargs)
|
||||
try:
|
||||
next(best_first_search(problem, graph=True)) # find 1st solution
|
||||
except InterruptedError:
|
||||
return None
|
||||
except StopIteration:
|
||||
print("Exhausted all possibilities without finding a complete solution ; or some irregularity occurred.")
|
||||
result = problem.get_solution_state()
|
||||
return result
|
||||
|
||||
|
||||
class WFC_GenParallel:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
gen_types = WFC_GenerateNode.NODE_INPUT_TYPES()
|
||||
gen_types["required"]["max_parallel_tasks"] = ("INT", {"default": 4, "min": 1, "max": 32})
|
||||
return gen_types
|
||||
|
||||
RETURN_TYPES = ("WFC_State",)
|
||||
RETURN_NAMES = ("state",)
|
||||
FUNCTION = "gen"
|
||||
CATEGORY = "Bmad/WFC"
|
||||
INPUT_IS_LIST = True
|
||||
OUTPUT_IS_LIST = (True,)
|
||||
|
||||
def gen(self, max_parallel_tasks, custom_temperature_config=None, custom_node_value_config=None, **kwargs):
|
||||
from multiprocessing import Manager
|
||||
from multiprocessing.managers import ValueProxy
|
||||
from joblib import Parallel, delayed
|
||||
from threading import Event, Thread
|
||||
|
||||
max_parallel_tasks = max_parallel_tasks[0]
|
||||
ct_len = 0 if custom_temperature_config is None else len(custom_temperature_config)
|
||||
cnv_len = 0 if custom_node_value_config is None else len(custom_node_value_config)
|
||||
|
||||
max_len = 0
|
||||
for v in kwargs.values():
|
||||
if (v_len := len(v)) > max_len:
|
||||
max_len = v_len
|
||||
|
||||
# TODO count is also done inside Problem, maybe should use as optional arg to avoid repeating the operation
|
||||
ss = kwargs["starting_state"]
|
||||
total_tiles_to_proc = sum([i.size-np.count_nonzero(i) for i in ss])
|
||||
total_tiles_to_proc += (ss[-1].size-np.count_nonzero(ss[-1]))*(max_len - len(ss))
|
||||
|
||||
manager = Manager()
|
||||
stop: ValueProxy = manager.Value('b', False) # Set to True to abort
|
||||
ticker: ValueProxy = manager.Value('i', 0) # counts max depth increments
|
||||
items = kwargs.items()
|
||||
per_gen_inputs = []
|
||||
for i in range(max_len):
|
||||
input_i = {item[0]: item[1][min(i, len(item[1]) - 1)] for item in items}
|
||||
if ct_len > 0:
|
||||
input_i.update(custom_temperature_config[min(i, ct_len - 1)])
|
||||
if cnv_len > 0:
|
||||
input_i.update(custom_node_value_config[min(i, cnv_len - 1)])
|
||||
input_i.update({"stop_proxy": stop})
|
||||
input_i.update({"ticker_proxy": ticker})
|
||||
per_gen_inputs.append(input_i)
|
||||
|
||||
finished_event = Event()
|
||||
pbar: utils.ProgressBar = utils.ProgressBar(total_tiles_to_proc)
|
||||
t = Thread(target=waiting_loop, args=(finished_event, stop, pbar, ticker, total_tiles_to_proc))
|
||||
t.start()
|
||||
|
||||
final_result = Parallel(n_jobs=max_parallel_tasks)(delayed(generate_single)(per_gen_inputs[i]) for i in range(max_len))
|
||||
|
||||
finished_event.set()
|
||||
if stop.get():
|
||||
throw_exception_if_processing_interrupted()
|
||||
|
||||
return (final_result,)
|
||||
@@ -0,0 +1 @@
|
||||
py_search
|
||||
@@ -0,0 +1,561 @@
|
||||
from collections import defaultdict
|
||||
from functools import lru_cache, cache
|
||||
from multiprocessing.managers import ValueProxy
|
||||
import numpy as np
|
||||
import cv2 as cv
|
||||
from numpy import ndarray
|
||||
from py_search.base import Problem, Node
|
||||
import hashlib
|
||||
|
||||
|
||||
TILE_DIGEST_SIZE = 4 # in bytes
|
||||
NP_ENCODED_TILE_TYPE = "longlong"
|
||||
WORLD_DIGEST_SIZE = 4
|
||||
|
||||
|
||||
class WFC_Sample:
|
||||
"""
|
||||
From a source image, compute & store the following data:
|
||||
tile_data : { tile hashcode : ( tile , frequency ) , ... }
|
||||
super_tile_data : [ ( 3x3 matrix of tile hashes, count ) ]
|
||||
tile_dims : ( tile height, width, channels )
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def tile_to_hash(tile):
|
||||
return int.from_bytes(hashlib.blake2b(tile.tobytes(), digest_size=TILE_DIGEST_SIZE).digest(), byteorder="big")
|
||||
|
||||
def __init__(self, src_imgs, cell_width, cell_height):
|
||||
self.tile_data, self.super_tile_data, self.tile_dims = self.prepare(src_imgs[0], cell_width, cell_height)
|
||||
|
||||
for img in src_imgs[1:]:
|
||||
r_tile_data, r_super_tile_data, _ = self.prepare(img, cell_width, cell_height)
|
||||
self.tile_data = {k: (v[0], self.tile_data.get(k, (None, 0))[1] + v[1])
|
||||
for k, v in {**self.tile_data, **r_tile_data}.items()}
|
||||
self.super_tile_data = self.merge_tuples(self.super_tile_data, r_super_tile_data)
|
||||
|
||||
@staticmethod
|
||||
def merge_tuples(list1, list2): # adapted from GPT; might be wrong
|
||||
result_dict = defaultdict(int)
|
||||
shape = list1[0][0].shape
|
||||
for ndarray, value in list1 + list2:
|
||||
# Convert ndarray to a hashable type (tuple)
|
||||
key = tuple(ndarray.flatten())
|
||||
result_dict[key] += value
|
||||
|
||||
# Convert the result back to a list of tuples
|
||||
result_list = [(np.array(key).reshape(shape), value) for key, value in result_dict.items()]
|
||||
return result_list
|
||||
|
||||
def get_tile_data(self) -> dict[int, tuple[ndarray, float]]:
|
||||
"""
|
||||
@return: hash to (tile, freq) pair dictionary
|
||||
"""
|
||||
return self.tile_data
|
||||
|
||||
def get_super_tile_data(self) -> list[tuple[ndarray, int]]:
|
||||
"""
|
||||
@return: list of (super tile, count) pairs
|
||||
"""
|
||||
return self.super_tile_data
|
||||
|
||||
@staticmethod
|
||||
def adjust_image_to_tile_size(img, tile_height, tile_width):
|
||||
"""
|
||||
@return: the adjusted image, height in number of cells, and width
|
||||
"""
|
||||
if img.shape[0] % tile_height != 0:
|
||||
print(f"src height is not divisible by cell_height!")
|
||||
if img.shape[1] % tile_width != 0:
|
||||
print(f"src width is not divisible by cell_height!")
|
||||
|
||||
height_in_tiles = round(img.shape[0] / tile_height)
|
||||
width_in_tiles = round(img.shape[1] / tile_width)
|
||||
new_height = tile_height * height_in_tiles
|
||||
new_width = tile_width * width_in_tiles
|
||||
|
||||
adjusted_image = cv.resize(img, (new_width, new_height)) if new_height != img.shape[0] or new_width != \
|
||||
img.shape[1] else img.copy()
|
||||
|
||||
return adjusted_image, height_in_tiles, width_in_tiles
|
||||
|
||||
@staticmethod
|
||||
def prepare(src_img, tile_width, tile_height):
|
||||
src_shape = src_img.shape
|
||||
adjusted_img, ycell_len, xcell_len = WFC_Sample.adjust_image_to_tile_size(src_img, tile_height, tile_width)
|
||||
tiles = adjusted_img.reshape(
|
||||
(
|
||||
ycell_len, # adjusted_img.shape[0] // tile_height,
|
||||
tile_height,
|
||||
xcell_len, # adjusted_img.shape[1] // tile_width,
|
||||
tile_width,
|
||||
src_shape[2]
|
||||
)
|
||||
).swapaxes(1, 2)
|
||||
size_in_tiles = tiles.shape[:2]
|
||||
tiles = tiles.reshape(-1, tile_height, tile_width, src_shape[2])
|
||||
utiles, counts = np.unique(tiles, axis=0, return_counts=True)
|
||||
ut_hashes = [WFC_Sample.tile_to_hash(tile) for tile in utiles]
|
||||
# assert len(set(ut_hashes)) == len(ut_hashes)
|
||||
|
||||
tiles_data = dict(zip(ut_hashes, zip(utiles, counts / tiles.shape[0])))
|
||||
hashed_tiles = np.array([WFC_Sample.tile_to_hash(tile) for tile in tiles])
|
||||
hashed_tiles = hashed_tiles.reshape(size_in_tiles)
|
||||
|
||||
super_tiles = np.array(
|
||||
[hashed_tiles[y:y + 3, x:x + 3] for y in range(size_in_tiles[0] - 2) for x in range(size_in_tiles[1] - 2)])
|
||||
|
||||
u_super_tiles, super_counts = np.unique(super_tiles, axis=0, return_counts=True)
|
||||
super_tiles_data = list(zip(u_super_tiles, super_counts))
|
||||
|
||||
return tiles_data, super_tiles_data, (tile_height, tile_width, src_shape[2])
|
||||
|
||||
def img_to_tile_encoded_world(self, src_img):
|
||||
adjusted_img, ycell_len, xcell_len = WFC_Sample.adjust_image_to_tile_size(src_img, *self.tile_dims[:2])
|
||||
tiles = adjusted_img.reshape(
|
||||
(
|
||||
ycell_len,
|
||||
self.tile_dims[0],
|
||||
xcell_len,
|
||||
self.tile_dims[1],
|
||||
adjusted_img.shape[2]
|
||||
)
|
||||
).swapaxes(1, 2)
|
||||
size_in_tiles = tiles.shape[:2]
|
||||
tiles = tiles.reshape(-1, self.tile_dims[0], self.tile_dims[1], adjusted_img.shape[2])
|
||||
hashed_world = [WFC_Sample.tile_to_hash(tile) for tile in tiles]
|
||||
hashed_world = np.array([tile if tile in self.tile_data.keys() else 0 for tile in hashed_world]).astype(
|
||||
NP_ENCODED_TILE_TYPE)
|
||||
return hashed_world.reshape(*size_in_tiles)
|
||||
|
||||
def tile_encoded_to_img(self, src_state: ndarray):
|
||||
th, tw = self.tile_dims[:2]
|
||||
img = np.zeros((src_state.shape[0] * th, src_state.shape[1] * tw, self.tile_dims[2]))
|
||||
mask = np.ones((src_state.shape[0] * th, src_state.shape[1] * tw)) * 255
|
||||
for (y, x), hashcode in np.ndenumerate(src_state):
|
||||
if hashcode in self.tile_data:
|
||||
img[y * th:(y + 1) * th, x * tw:(x + 1) * tw] = self.tile_data[hashcode][0]
|
||||
mask[y * th:(y + 1) * th, x * tw:(x + 1) * tw] = np.zeros(self.tile_data[hashcode][0].shape[:2])
|
||||
return img, mask
|
||||
|
||||
|
||||
class WFC_Problem(Problem):
|
||||
def __init__(self, sample: WFC_Sample, starting_state: ndarray, seed: int = 0, use_8_cardinals: bool = False,
|
||||
max_freq_adjust: float = 1, plateau_check_interval: int = -1,
|
||||
starting_temperature: float = 50, min_min_temperature: float = 0, max_min_temperature: float = 80,
|
||||
reverse_depth_w: float = 1, node_cost_w: float = 1, path_entropy_average_w: float = 0,
|
||||
stop_proxy: ValueProxy = None, ticker_proxy: ValueProxy = None
|
||||
):
|
||||
"""
|
||||
@param sample: contains the tiles, their frequencies and "implicit" constraints
|
||||
@param starting_state: complete the provided state instead of starting with an empty world.
|
||||
If provided, width and height are ignored.
|
||||
@param seed: used to setup the generator so that the result can be reproducible ( deterministic )
|
||||
@param use_8_cardinals: consider the surrounding 8 tiles if set to TRUE; OR only the 4 cardinals if set to FALSE
|
||||
@param max_freq_adjust: scale frequency adjustment weight.
|
||||
if set to zero, the frequency of the tiles registered in the given samples is ignored.
|
||||
@param plateau_check_interval: the number of nodes to be processed before checking the highest registered depth.
|
||||
if the depth hasn't changed between checks, the search is stopped.
|
||||
Set to 0 (zero) to ignore plateau checks.
|
||||
Set to -1 (minus 1) to auto select depending on the number of nodes to process.
|
||||
@param starting_temperature:
|
||||
@param min_min_temperature:
|
||||
@param max_min_temperature:
|
||||
@param reverse_depth_w:
|
||||
@param node_cost_w:
|
||||
@param path_entropy_average_w:
|
||||
"""
|
||||
super().__init__(initial=(0, 0), initial_cost=0, extra=0)
|
||||
# initial state -> (depth, hash) -> 0 represents empty world at the start, at zero depth
|
||||
# extra -> the sum of entropies of the closed nodes in the current branch ( at the moment they were closed )
|
||||
|
||||
# BASIC DATA
|
||||
self._stop_proxy = stop_proxy
|
||||
self._ticker_proxy = ticker_proxy
|
||||
self.rng = np.random.default_rng(seed=seed)
|
||||
self._sample: WFC_Sample = sample
|
||||
|
||||
non_zeroes = np.count_nonzero(starting_state)
|
||||
self._number_of_tiles_to_process = starting_state.size - non_zeroes
|
||||
self._world_tdims = starting_state.shape
|
||||
self._temp_world_state = starting_state.copy()
|
||||
self._starting_state = starting_state.copy() if non_zeroes > 0 else None
|
||||
# _starting_state has 2 internal uses:
|
||||
# 1. if None the center tile is set to open, otherwise the state is iterated to find the tiles at the edges
|
||||
# 2. initialize state to return instead of reverting last node actions
|
||||
|
||||
# KEEP TRACK OF OPEN TILES ( yet to explore after the last closed node )
|
||||
self._temp_world_open_super_tiles: set[tuple[int, int]] = set([])
|
||||
|
||||
# INFLUENCE COST WEIGHTS & FINAL NODE VALUE
|
||||
# influences the nodes' costs. high temperature lowers the influence of random noise and frequency adjustments
|
||||
self._min_temperature = starting_temperature
|
||||
self._min_min_temperature, self._max_min_temperature = min_min_temperature, max_min_temperature
|
||||
self._rev_depth_w, self._node_cost_w, self._path_ent_avg_w = reverse_depth_w, node_cost_w, path_entropy_average_w
|
||||
|
||||
# STOP THE SEARCH
|
||||
self._best_node = None
|
||||
self._prev_best_depth = 0
|
||||
self._last_node = None
|
||||
self._plateau_check_ticker = 0
|
||||
self._plateau_stop_steps = self._world_tdims[0] * self._world_tdims[1] / 2.0 \
|
||||
if plateau_check_interval == -1 else plateau_check_interval
|
||||
|
||||
self._stop: bool = False # stops the search; set to True when all the tiles are filled OR a plateau is reached
|
||||
|
||||
# OTHERS
|
||||
self._use_8cardinals = use_8_cardinals
|
||||
|
||||
tile_data = sample.get_tile_data()
|
||||
self._tile_counts = dict(zip(tile_data.keys(), [0] * len(tile_data)))
|
||||
|
||||
self._max_freq_adjust = max_freq_adjust
|
||||
t = self._number_of_tiles_to_process
|
||||
a = [[0, 0, 1], [t ** 2, t, 1], [(t / 2) ** 2, t / 2, 1]]
|
||||
b = [t, t, 0]
|
||||
self._tile_freq_adjustment_poly = np.poly1d(np.linalg.solve(a, b))
|
||||
|
||||
@lru_cache(maxsize=8)
|
||||
def temp_ratio(self, node_depth: int, prior_node_depth: int):
|
||||
# TODO -> potentially something to change/customize
|
||||
depth_diff = prior_node_depth - node_depth
|
||||
depth_ratio = node_depth / self._number_of_tiles_to_process
|
||||
ratio = min(.9, (abs(depth_diff) * 3) / np.sqrt(self._number_of_tiles_to_process)) if depth_diff != 0 else \
|
||||
depth_ratio ** 2.5 / 80
|
||||
return ratio
|
||||
|
||||
def get_new_temperature(self, node_depth: int, prior_node_depth: int):
|
||||
limit = self._max_min_temperature if node_depth <= prior_node_depth else self._min_min_temperature
|
||||
ratio = self.temp_ratio(node_depth, prior_node_depth)
|
||||
return limit * ratio + self._min_temperature * (1 - ratio)
|
||||
|
||||
@lru_cache(maxsize=32)
|
||||
def _tile_freq_adjustment_func(self, depth):
|
||||
return self._max_freq_adjust * (1 - self._tile_freq_adjustment_poly(depth) / self._number_of_tiles_to_process)
|
||||
|
||||
def update_open_nodes(self, y, x, is_reopening):
|
||||
"""
|
||||
Update open nodes when updating world state
|
||||
@param is_reopening:
|
||||
@return:
|
||||
"""
|
||||
# 0 => nothing; 1 => open; 2 => closed
|
||||
wost = self._temp_world_open_super_tiles
|
||||
|
||||
indices_to_check = [(y - 1 + _y, x - 1 + _x) for _y in range(3) for _x in range(3) if _y != 1 or _x != 1]
|
||||
if not self._use_8cardinals:
|
||||
indices_to_check = [indices_to_check[i] for i in [1, 3, 4, 6]]
|
||||
indices_to_check = [(_y, _x) for (_y, _x) in indices_to_check
|
||||
if 0 <= _y < self._temp_world_state.shape[0] and 0 <= _x < self._temp_world_state.shape[1]]
|
||||
|
||||
if is_reopening:
|
||||
wost.add((y, x))
|
||||
|
||||
for (_y, _x) in indices_to_check:
|
||||
if not wost.__contains__((_y, _x)):
|
||||
continue
|
||||
|
||||
sub_indices_to_check = [(_y - 1 + __y, _x - 1 + __x) for __y in range(3)
|
||||
for __x in range(3)]
|
||||
if not self._use_8cardinals:
|
||||
sub_indices_to_check = [sub_indices_to_check[i] for i in [1, 3, 5, 7]]
|
||||
sub_indices_to_check = [(__y, __x) for (__y, __x) in sub_indices_to_check
|
||||
if 0 <= __y < self._temp_world_state.shape[0] and 0 <= __x <
|
||||
self._temp_world_state.shape[1]]
|
||||
|
||||
if any(self._temp_world_state[__y, __x] != 0 for (__y, __x) in sub_indices_to_check):
|
||||
continue # -> position should remain open
|
||||
# otherwise -> position should be closed
|
||||
wost.remove((_y, _x))
|
||||
|
||||
return
|
||||
|
||||
# otherwise -> closing
|
||||
wost.remove((y, x))
|
||||
for _y, _x in indices_to_check:
|
||||
if self._temp_world_state[_y, _x] == 0:
|
||||
wost.add((_y, _x))
|
||||
|
||||
def get_cell_potential_states_and_costs(self, y, x, world_state, depth) -> \
|
||||
tuple[list[bytes], ndarray | None, float | None]:
|
||||
world_indices_to_check = [(y - 1 + _y, x - 1 + _x) for _y in range(3) for _x in range(3) if _y != 1 or _x != 1]
|
||||
adjacent_states = [world_state[_y, _x] if 0 <= _y < world_state.shape[0] and 0 <= _x < world_state.shape[1]
|
||||
else 0 for (_y, _x) in world_indices_to_check]
|
||||
tiles, probabilities = \
|
||||
self.get_cell_potential_states_8cardinals(*adjacent_states) \
|
||||
if self._use_8cardinals else \
|
||||
self.get_cell_potential_states_4cardinals(*[adjacent_states[i] for i in [1, 3, 4, 6]])
|
||||
|
||||
if len(tiles) == 0:
|
||||
return [], None, None # nothing to compute, so just return early
|
||||
|
||||
# generate random weights and temperature
|
||||
rands = self.rng.random(len(probabilities))
|
||||
temp = 1 - self.rng.integers(low=int(self._min_temperature), high=100) / 100.0
|
||||
|
||||
# GET tile type freq in original samples AND in current generation
|
||||
tile_data = self._sample.get_tile_data()
|
||||
sample_freqs = np.array([tile_data[t][1] for t in tiles])
|
||||
current_counts = np.array([self._tile_counts[t] for t in tiles])
|
||||
current_freqs = current_counts / max(1, depth)
|
||||
|
||||
if self._max_freq_adjust == 0.0:
|
||||
adjusted_freqs_diff = np.zeros(probabilities.size)
|
||||
else:
|
||||
depth_adjustment = self._tile_freq_adjustment_func(depth)
|
||||
adjusted_freqs_diff = np.sign(sample_freqs - current_freqs) * (
|
||||
1 - np.minimum(sample_freqs, current_freqs) / np.maximum(sample_freqs, current_freqs))
|
||||
adjusted_freqs_diff *= depth_adjustment
|
||||
# ===========================================================================
|
||||
entropy: float = - np.sum(probabilities * np.log2(probabilities)) / np.log2(len(probabilities)) if len(
|
||||
probabilities) > 1 else 0
|
||||
|
||||
# TODO -> review formula... multiplication w/ entropy may be too strong for some rules;
|
||||
# consider attenuating entropy on low temperatures or when adjusting freqs
|
||||
# OR, consider some node to customize the costs later
|
||||
costs = (1 - np.clip(probabilities
|
||||
+ adjusted_freqs_diff * temp
|
||||
+ (rands * 2 - 1) * temp
|
||||
, 0, 1)) * entropy
|
||||
|
||||
return tiles, costs, entropy
|
||||
|
||||
@cache
|
||||
def get_cell_potential_states_8cardinals(self, p1, p2, p3, p4, p6, p7, p8, p9):
|
||||
"""
|
||||
the state of adjacent cells; where 0 = unknown
|
||||
[[1,2,3],
|
||||
[4,c,6],
|
||||
[7,8,9]]
|
||||
@return: list of tuples (state, prob.)
|
||||
"""
|
||||
|
||||
def is_possible(super_tile):
|
||||
tiles = np.array([[p1, p2, p3], [p4, 0, p6], [p7, p8, p9]])
|
||||
for (y, x), tile in np.ndenumerate(tiles):
|
||||
# print(f"super={super_tile} ; index = {(y,x)}")
|
||||
if tile == 0:
|
||||
continue
|
||||
if super_tile[y, x] != tile:
|
||||
return False
|
||||
return True
|
||||
|
||||
# filter all the possible super tiles
|
||||
pcs = defaultdict(int)
|
||||
for stile, count in self._sample.get_super_tile_data():
|
||||
if is_possible(stile):
|
||||
pcs[stile[1, 1]] += count
|
||||
|
||||
return self.map_to_probabilities(pcs)
|
||||
|
||||
@cache
|
||||
def get_cell_potential_states_4cardinals(self, p2, p4, p6, p8):
|
||||
"""
|
||||
Read get_cell_potential_states_8cardinals documentation.
|
||||
This function is similar, but only takes into account 4 cardinals
|
||||
"""
|
||||
|
||||
def is_possible(super_tile):
|
||||
tiles = np.array([[0, p2, 0], [p4, 0, p6], [0, p8, 0]])
|
||||
for (y, x), tile in np.ndenumerate(tiles):
|
||||
if tile == 0:
|
||||
continue
|
||||
if super_tile[y, x] != tile:
|
||||
return False
|
||||
return True
|
||||
|
||||
# filter all the possible super tiles
|
||||
pcs = defaultdict(int)
|
||||
for stile, count in self._sample.get_super_tile_data():
|
||||
if is_possible(stile):
|
||||
pcs[stile[1, 1]] += count
|
||||
|
||||
return self.map_to_probabilities(pcs)
|
||||
|
||||
def map_to_probabilities(self, pcs):
|
||||
if len(pcs) == 0:
|
||||
return [], []
|
||||
# compute probabilities for each possible state
|
||||
counts = list(pcs.values())
|
||||
total_counts: float = sum(counts)
|
||||
probabilities: ndarray = np.array(counts).astype(np.float32) / total_counts
|
||||
assert np.sum(probabilities > 1) == 0 # seems fine
|
||||
return pcs.keys(), probabilities
|
||||
|
||||
def node_value(self, node: Node):
|
||||
"""
|
||||
The function used to compute the value of a node.
|
||||
"""
|
||||
rev_depth = (1 + self._number_of_tiles_to_process - node.depth())
|
||||
# can be used to prioritize nodes w/ high depth, for a quicker generation
|
||||
|
||||
node_cost = node.cost()
|
||||
# if temperature is high, this is the most promising locally
|
||||
|
||||
path_avg_entropy = node.extra / (1 + node.depth())
|
||||
# can be used to backtrack preemptively if path doesn't look strong
|
||||
# the entropies used are the ones obtain when closing a node, so this might be misleading indicator
|
||||
|
||||
return rev_depth * self._rev_depth_w + node_cost * self._node_cost_w + path_avg_entropy * self._path_ent_avg_w
|
||||
|
||||
def successors(self, node):
|
||||
from py_search.base import Node
|
||||
"""
|
||||
Generate all possible next states
|
||||
"""
|
||||
if (self._stop_proxy is not None and self._stop_proxy.get()) :
|
||||
#print("")
|
||||
raise InterruptedError()
|
||||
|
||||
# world state is kept in self._temp_world_state; updated here when closing the node.
|
||||
# get_world_state func rollbacks any actions when depth is maintained or decreased.
|
||||
|
||||
cum_entropy = 0
|
||||
if node.depth() == 0:
|
||||
self._best_node = node
|
||||
self.open_nodes_on_depth_zero()
|
||||
self._last_node = node
|
||||
world_state = self._temp_world_state
|
||||
else:
|
||||
world_state = self.get_world_state(self._last_node, node)
|
||||
self._min_temperature = self.get_new_temperature(node.depth(), self._last_node.depth())
|
||||
self._last_node = node
|
||||
cum_entropy = node.parent.extra
|
||||
|
||||
depth = node.depth()
|
||||
if depth > self._best_node.depth():
|
||||
self._best_node = node
|
||||
if self._ticker_proxy is not None:
|
||||
self._ticker_proxy.set(self._ticker_proxy.get()+1)
|
||||
if depth >= self._number_of_tiles_to_process:
|
||||
self._stop = True
|
||||
print("\nEnded search with all tiles filled.")
|
||||
return
|
||||
|
||||
if self._plateau_stop_steps > 0:
|
||||
self._plateau_check_ticker += 1
|
||||
if self._plateau_check_ticker >= self._plateau_stop_steps:
|
||||
if self._prev_best_depth == self._best_node.depth():
|
||||
self._stop = True
|
||||
print("\nEnded due to depth plateauing.")
|
||||
if self._ticker_proxy is not None:
|
||||
self._ticker_proxy.set(self._ticker_proxy.get() + self._number_of_tiles_to_process - self._best_node.depth())
|
||||
return
|
||||
self._plateau_check_ticker = 0
|
||||
self._prev_best_depth = self._best_node.depth()
|
||||
|
||||
iyxs = self._temp_world_open_super_tiles
|
||||
potential_collapses = [self.get_cell_potential_states_and_costs(y, x, world_state, depth) for (y, x) in iyxs]
|
||||
|
||||
#print(f"depth = {depth:5,.0f} | temperature={self._min_temperature:5,.1f} | "
|
||||
# f"freq_depth_adjustment={self._tile_freq_adjustment_func(depth):6,.2f} | "
|
||||
# f"open tiles:{len(iyxs):5,.0f} ", end="\r")
|
||||
|
||||
if len(potential_collapses) == 0 or any(counts is None for (_, counts, _) in potential_collapses):
|
||||
# this state is impossible, so don't return any of its children
|
||||
node.node_cost = float("inf") # the node has now been closed, but it could help w/ debugging
|
||||
return
|
||||
|
||||
for (y, x), (potential_states, costs, entropy) in zip(iyxs, potential_collapses):
|
||||
items = zip(potential_states, costs)
|
||||
|
||||
# if self._temperature > xxx: # TODO consider making this an option set by the user
|
||||
# # items = sorted(items, key=itemgetter(1))[:2]
|
||||
# items = sorted(items, key=itemgetter(1))
|
||||
# take = round(2 * self._temperature / self._max_temperature + len(items) * (
|
||||
# 1 - self._temperature / self._max_temperature))
|
||||
# items = items[:take]
|
||||
|
||||
for tile_type, cost in items:
|
||||
world_state[y, x] = tile_type
|
||||
world_state_hash = int.from_bytes(hashlib.blake2b(world_state.tobytes(), digest_size=4).digest(),
|
||||
byteorder="big")
|
||||
world_state[y, x] = 0
|
||||
yield Node(state=(depth + 1, world_state_hash),
|
||||
parent=node, action=((y, x), tile_type), node_cost=cost, extra=cum_entropy + entropy)
|
||||
|
||||
def goal_test(self, state_node, goal_node=None):
|
||||
# state is not kept in each node, so the checks are done when closing a node.
|
||||
# goal_test is only defined to terminate the search;
|
||||
# the real goal test is done in the successors method
|
||||
return self._stop
|
||||
|
||||
def get_world_state(self, last_node: Node, current_node: Node):
|
||||
start_depth = min(last_node.depth(), current_node.depth())
|
||||
|
||||
p_node = last_node
|
||||
c_node = current_node
|
||||
|
||||
for _ in range(p_node.depth() - start_depth):
|
||||
self.revert_action(p_node.action)
|
||||
p_node = p_node.parent
|
||||
|
||||
for _ in range(c_node.depth() - start_depth):
|
||||
c_node = c_node.parent
|
||||
|
||||
common_depth = 0
|
||||
for i in range(start_depth + 1)[::-1]:
|
||||
if p_node.state == c_node.state:
|
||||
common_depth = i
|
||||
break
|
||||
self.revert_action(p_node.action)
|
||||
p_node = p_node.parent
|
||||
c_node = c_node.parent
|
||||
|
||||
# NOTE: node is removed from opened set when closing; thus, the order of operations needs to be preserved
|
||||
c_node = current_node
|
||||
nodes_to_apply = []
|
||||
for _ in range(current_node.depth() - common_depth):
|
||||
nodes_to_apply.append(c_node)
|
||||
c_node = c_node.parent
|
||||
|
||||
for node in nodes_to_apply[::-1]:
|
||||
self.apply_action(node.action)
|
||||
|
||||
return self._temp_world_state
|
||||
|
||||
def open_nodes_on_depth_zero(self):
|
||||
if self._starting_state is None:
|
||||
# open center tile
|
||||
self._temp_world_open_super_tiles.add((self._world_tdims[0] // 2, self._world_tdims[1] // 2))
|
||||
return
|
||||
# otherwise -> find all in starting state
|
||||
|
||||
relative_indices_to_check = [(_y, _x) for _y in range(3) for _x in range(3) if _y != 1 or _x != 1]
|
||||
if not self._use_8cardinals:
|
||||
relative_indices_to_check = [relative_indices_to_check[i] for i in [1, 3, 4, 6]]
|
||||
|
||||
wost = self._temp_world_open_super_tiles
|
||||
for (y, x), tile in np.ndenumerate(self._starting_state):
|
||||
if tile != 0:
|
||||
continue
|
||||
indices_to_check = [(__y, __x) for (_y, _x) in relative_indices_to_check
|
||||
if 0 <= (__y := y + _y) < self._temp_world_state.shape[0]
|
||||
and 0 <= (__x := x + _x) < self._temp_world_state.shape[1]]
|
||||
|
||||
if any(self._starting_state[__y, __x] != 0 for (__y, __x) in indices_to_check):
|
||||
wost.add((y, x))
|
||||
|
||||
def revert_action(self, node_action):
|
||||
(pos, state) = node_action
|
||||
self._tile_counts[self._temp_world_state[*pos]] -= 1
|
||||
self._temp_world_state[*pos] = 0
|
||||
self.update_open_nodes(*pos, is_reopening=True)
|
||||
|
||||
def apply_action(self, node_action):
|
||||
(pos, state) = node_action
|
||||
self._temp_world_state[*pos] = state
|
||||
self._tile_counts[state] += 1
|
||||
self.update_open_nodes(*pos, is_reopening=False)
|
||||
|
||||
def get_solution_state(self):
|
||||
node: Node = self._best_node
|
||||
encoded_state = np.zeros(self._temp_world_state.shape[:2]) if self._starting_state is None \
|
||||
else self._starting_state.copy()
|
||||
for d in range(node.depth()):
|
||||
(y, x), tile_hash = node.action
|
||||
node = node.parent
|
||||
if tile_hash != 0:
|
||||
encoded_state[y, x] = tile_hash
|
||||
|
||||
return encoded_state.astype(NP_ENCODED_TILE_TYPE)
|
||||
Reference in New Issue
Block a user