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)" }