comfyui-model-pack / custom_nodes /NeuralISP /neural_isp_node.py
45th's picture
Upload 91 files
0e7beee verified
Raw History Blame Contribute Delete
7.04 kB
import os
import torch
import torch.nn.functional as F
import folder_paths
from .mirnetv2_model import MIRNet_v2
MODEL_DIR = os.path.join(
folder_paths.models_dir,
"mirnetv2"
)
class NeuralISPNode:
@classmethod
def INPUT_TYPES(cls):
models = []
if os.path.exists(MODEL_DIR):
for f in os.listdir(MODEL_DIR):
if f.endswith((".pth", ".pt")):
models.append(f)
if not models:
models = [
"enhancement_fivek.pth",
"enhancement_lol.pth",
"real_denoising.pth"
]
return {
"required": {
"image":
("IMAGE",),
"model":
(models,),
"strength":
(
"FLOAT",
{
"default":0.2,
"min":0.0,
"max":1.0,
"step":0.05
}
),
"tile_size":
(
"INT",
{
"default":512,
"min":128,
"max":2048,
"step":64
}
),
"overlap":
(
"INT",
{
"default":128,
"min":32,
"max":512,
"step":32
}
)
}
}
RETURN_TYPES = ("IMAGE",)
FUNCTION = "apply"
CATEGORY = "📸 shots by Faded/Neural Post FX"
def load_model(self, filename):
device = (
"cuda"
if torch.cuda.is_available()
else "cpu"
)
model = MIRNet_v2()
path = os.path.join(
MODEL_DIR,
filename
)
checkpoint = torch.load(
path,
map_location="cpu"
)
if "params" in checkpoint:
checkpoint = checkpoint["params"]
elif "state_dict" in checkpoint:
checkpoint = checkpoint["state_dict"]
elif "model" in checkpoint:
checkpoint = checkpoint["model"]
cleaned = {}
for k,v in checkpoint.items():
if k.startswith("module."):
k = k[7:]
cleaned[k] = v
model.load_state_dict(
cleaned,
strict=False
)
model.eval()
model.to(device)
# MIRNet лучше держать FP32
model.float()
return model, device
def create_mask(
self,
h,
w,
device
):
mask = torch.ones(
1,
1,
h,
w,
device=device
)
return mask
def process_tile(
self,
model,
img,
tile_size,
overlap
):
_,_,h,w = img.shape
stride = tile_size - overlap
output = torch.zeros_like(
img,
dtype=torch.float32
)
weight = torch.zeros_like(
img,
dtype=torch.float32
)
for y in range(0,h,stride):
for x in range(0,w,stride):
y1 = min(
y+tile_size,
h
)
x1 = min(
x+tile_size,
w
)
y0 = max(
0,
y1-tile_size
)
x0 = max(
0,
x1-tile_size
)
patch = img[
:,
:,
y0:y1,
x0:x1
]
ph = patch.shape[-2]
pw = patch.shape[-1]
# MIRNet требует размеры кратные 32
pad_h = (
32 - ph % 32
) % 32
pad_w = (
32 - pw % 32
) % 32
patch_pad = F.pad(
patch,
(
0,
pad_w,
0,
pad_h
),
mode="reflect"
)
with torch.no_grad():
result = model(
patch_pad.float()
)
# возвращаем исходный размер
result = result[
:,
:,
:ph,
:pw
]
mask = self.create_mask(
ph,
pw,
img.device
)
output[
:,
:,
y0:y1,
x0:x1
] += (
result *
mask
)
weight[
:,
:,
y0:y1,
x0:x1
] += mask
return (
output /
weight.clamp(min=1e-6)
)
def apply(
self,
image,
model,
strength,
tile_size,
overlap
):
net,device = self.load_model(
model
)
img = image[0]
img = img.permute(
2,
0,
1
).unsqueeze(0)
img = img.to(device)
with torch.no_grad():
result = self.process_tile(
net,
img,
tile_size,
overlap
)
result = (
result * strength
+
img.float() *
(1-strength)
)
result = result.clamp(
0,
1
)
result = result.squeeze(0)
result = result.permute(
1,
2,
0
)
return (
result.unsqueeze(0)
.cpu(),
)
NODE_CLASS_MAPPINGS = {
"NeuralISP":
NeuralISPNode
}
NODE_DISPLAY_NAME_MAPPINGS = {
"NeuralISP":
"📸Neural ISP (x MIRNet)"
}