Refactor BLIP! Some some other bugs fixed as well.
This commit is contained in:
+79
-154
@@ -200,8 +200,6 @@ was_conf_template = {
|
||||
"text_nodes_type": "STRING",
|
||||
"webui_styles": None,
|
||||
"webui_styles_persistent_update": True,
|
||||
"blip_model_url": "https://storage.googleapis.com/sfr-vision-language-research/BLIP/models/model_base_capfilt_large.pth",
|
||||
"blip_model_vqa_url": "https://storage.googleapis.com/sfr-vision-language-research/BLIP/models/model_base_vqa_capfilt_large.pth",
|
||||
"sam_model_vith_url": "https://dl.fbaipublicfiles.com/segment_anything/sam_vit_h_4b8939.pth",
|
||||
"sam_model_vitl_url": "https://dl.fbaipublicfiles.com/segment_anything/sam_vit_l_0b3195.pth",
|
||||
"sam_model_vitb_url": "https://dl.fbaipublicfiles.com/segment_anything/sam_vit_b_01ec64.pth",
|
||||
@@ -1924,9 +1922,7 @@ class WAS_Tools_Class():
|
||||
img = img.filter(ImageFilter.GaussianBlur(radius=blur))
|
||||
|
||||
return img
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
# Version 2 optimized based on Mark Setchell's ideas
|
||||
def gradient_map(self, image, gradient_map, reverse=False):
|
||||
@@ -2398,9 +2394,31 @@ class WAS_Tools_Class():
|
||||
|
||||
return palette, '\n'.join(hex_palette)
|
||||
|
||||
#! IMAGE FILTER NODES
|
||||
|
||||
# IMAGE ADJUSTMENTS NODES
|
||||
from transformers import BlipProcessor, BlipForConditionalGeneration, BlipForQuestionAnswering
|
||||
|
||||
class BlipWrapper:
|
||||
def __init__(self, caption_model_id="Salesforce/blip-image-captioning-base", vqa_model_id="Salesforce/blip-vqa-base", device="cuda", cache_dir=None):
|
||||
self.device = torch.device(device='cuda' if device == "cuda" and torch.cuda.is_available() else 'cpu')
|
||||
self.caption_processor = BlipProcessor.from_pretrained(caption_model_id, cache_dir=cache_dir)
|
||||
self.caption_model = BlipForConditionalGeneration.from_pretrained(caption_model_id, cache_dir=cache_dir).to(self.device)
|
||||
self.vqa_processor = BlipProcessor.from_pretrained(vqa_model_id, cache_dir=cache_dir)
|
||||
self.vqa_model = BlipForQuestionAnswering.from_pretrained(vqa_model_id, cache_dir=cache_dir).to(self.device)
|
||||
|
||||
def generate_caption(self, image: Image.Image, min_length=50, max_length=100, num_beams=5, no_repeat_ngram_size=2, early_stopping=False):
|
||||
self.caption_model.eval()
|
||||
inputs = self.caption_processor(images=image, return_tensors="pt").to(self.device)
|
||||
outputs = self.caption_model.generate(**inputs, min_length=min_length, max_length=max_length, num_beams=num_beams, no_repeat_ngram_size=no_repeat_ngram_size, early_stopping=early_stopping)
|
||||
return self.caption_processor.decode(outputs[0], skip_special_tokens=True)
|
||||
|
||||
def answer_question(self, image: Image.Image, question: str, min_length=50, max_length=100, num_beams=5, no_repeat_ngram_size=2, early_stopping=False):
|
||||
self.vqa_model.eval()
|
||||
inputs = self.vqa_processor(images=image, text=question, return_tensors="pt").to(self.device)
|
||||
answer_ids = self.vqa_model.generate(**inputs, min_length=min_length, max_length=max_length, num_beams=num_beams, no_repeat_ngram_size=no_repeat_ngram_size, early_stopping=early_stopping)
|
||||
return self.vqa_processor.decode(answer_ids[0], skip_special_tokens=True)
|
||||
|
||||
|
||||
#! IMAGE FILTER NODES
|
||||
|
||||
# IMAGE SHADOW AND HIGHLIGHT ADJUSTMENTS
|
||||
|
||||
@@ -2439,7 +2457,7 @@ class WAS_Shadow_And_Highlight_Adjustment:
|
||||
return (pil2tensor(result), pil2tensor(shadows), pil2tensor(highlights) )
|
||||
|
||||
|
||||
# IMAGE PIXATE
|
||||
# IMAGE PIXELATE
|
||||
|
||||
class WAS_Image_Pixelate:
|
||||
def __init__(self):
|
||||
@@ -9481,7 +9499,7 @@ class WAS_Text_Multiline:
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"text": ("STRING", {"default": '', "multiline": True}),
|
||||
"text": ("STRING", {"default": '', "multiline": True, "dynamicPrompts": False}),
|
||||
}
|
||||
}
|
||||
RETURN_TYPES = (TEXT_TYPE,)
|
||||
@@ -10311,7 +10329,7 @@ class WAS_Text_Save:
|
||||
"path": ("STRING", {"default": './ComfyUI/output/[time(%Y-%m-%d)]', "multiline": False}),
|
||||
"filename_prefix": ("STRING", {"default": "ComfyUI"}),
|
||||
"filename_delimiter": ("STRING", {"default":"_"}),
|
||||
"filename_number_padding": ("INT", {"default":4, "min":2, "max":9, "step":1}),
|
||||
"filename_number_padding": ("INT", {"default":4, "min":0, "max":9, "step":1}),
|
||||
|
||||
}
|
||||
}
|
||||
@@ -10350,7 +10368,8 @@ class WAS_Text_Save:
|
||||
return (text, { "ui": { "string": text } } )
|
||||
|
||||
def generate_filename(self, path, prefix, delimiter, number_padding, extension):
|
||||
pattern = f"{re.escape(prefix)}{re.escape(delimiter)}(\\d{{{number_padding}}})"
|
||||
# Adjust the regex pattern to match one or more digits without fixed padding
|
||||
pattern = f"{re.escape(prefix)}{re.escape(delimiter)}(\\d+)"
|
||||
existing_counters = [
|
||||
int(re.search(pattern, filename).group(1))
|
||||
for filename in os.listdir(path)
|
||||
@@ -10363,10 +10382,19 @@ class WAS_Text_Save:
|
||||
else:
|
||||
counter = 1
|
||||
|
||||
filename = f"{prefix}{delimiter}{counter:0{number_padding}}{extension}"
|
||||
# Adjust the format string based on number_padding
|
||||
if number_padding > 0:
|
||||
filename = f"{prefix}{delimiter}{counter:0{number_padding}}{extension}"
|
||||
else:
|
||||
filename = f"{prefix}{delimiter}{counter}{extension}"
|
||||
|
||||
# Ensure the generated filename does not already exist
|
||||
while os.path.exists(os.path.join(path, filename)):
|
||||
counter += 1
|
||||
filename = f"{prefix}{delimiter}{counter:0{number_padding}}{extension}"
|
||||
if number_padding > 0:
|
||||
filename = f"{prefix}{delimiter}{counter:0{number_padding}}{extension}"
|
||||
else:
|
||||
filename = f"{prefix}{delimiter}{counter}{extension}"
|
||||
|
||||
return filename
|
||||
|
||||
@@ -10952,7 +10980,9 @@ class WAS_BLIP_Model_Loader:
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"blip_model": (["caption", "interrogate"], ),
|
||||
"blip_model": ("STRING", {"default": "Salesforce/blip-image-captioning-base"}),
|
||||
"vqa_model_id": ("STRING", {"default": "Salesforce/blip-vqa-base"}),
|
||||
"device": (["cuda", "cpu"],),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -10961,64 +10991,17 @@ class WAS_BLIP_Model_Loader:
|
||||
|
||||
CATEGORY = "WAS Suite/Loaders"
|
||||
|
||||
def blip_model(self, blip_model):
|
||||
def blip_model(self, blip_model, vqa_model_id, device):
|
||||
|
||||
if ( 'timm' not in packages()
|
||||
or 'transformers' not in packages()
|
||||
or 'fairscale' not in packages() ):
|
||||
cstr(f"Modules or packages are missing to use BLIP models. Please run the `{os.path.join(WAS_SUITE_ROOT, 'requirements.txt')}` through ComfyUI's python executable.").error.print()
|
||||
exit
|
||||
blip_dir = os.path.join(comfy_paths.models_dir, "blip")
|
||||
|
||||
if 'transformers==4.26.1' not in packages(True):
|
||||
cstr(f"`transformers==4.26.1` is required for BLIP models. Please run the `{os.path.join(WAS_SUITE_ROOT, 'requirements.txt')}` through ComfyUI's python executable.").error.print()
|
||||
exit
|
||||
# Attempt legacy support
|
||||
if blip_model in ("caption", "interrogate"):
|
||||
blip_model = "Salesforce/blip-image-captioning-base"
|
||||
|
||||
device = 'cpu'
|
||||
conf = getSuiteConfig()
|
||||
size = 384
|
||||
|
||||
if blip_model == 'caption':
|
||||
|
||||
from .modules.BLIP.blip_module import blip_decoder
|
||||
|
||||
blip_dir = os.path.join(MODELS_DIR, 'blip')
|
||||
if not os.path.exists(blip_dir):
|
||||
os.makedirs(blip_dir, exist_ok=True)
|
||||
|
||||
torch.hub.set_dir(blip_dir)
|
||||
|
||||
if conf.__contains__('blip_model_url'):
|
||||
model_url = conf['blip_model_url']
|
||||
else:
|
||||
model_url = 'https://storage.googleapis.com/sfr-vision-language-research/BLIP/models/model_base_capfilt_large.pth'
|
||||
|
||||
model = blip_decoder(pretrained=model_url, image_size=size, vit='base')
|
||||
model.eval()
|
||||
model = model.to(device)
|
||||
|
||||
elif blip_model == 'interrogate':
|
||||
|
||||
from .modules.BLIP.blip_module import blip_vqa
|
||||
|
||||
blip_dir = os.path.join(MODELS_DIR, 'blip')
|
||||
if not os.path.exists(blip_dir):
|
||||
os.makedirs(blip_dir, exist_ok=True)
|
||||
|
||||
torch.hub.set_dir(blip_dir)
|
||||
|
||||
if conf.__contains__('blip_model_vqa_url'):
|
||||
model_url = conf['blip_model_vqa_url']
|
||||
else:
|
||||
model_url = 'https://storage.googleapis.com/sfr-vision-language-research/BLIP/models/model_base_vqa_capfilt_large.pth'
|
||||
|
||||
model = blip_vqa(pretrained=model_url, image_size=size, vit='base')
|
||||
model.eval()
|
||||
model = model.to(device)
|
||||
|
||||
result = ( model, blip_model )
|
||||
|
||||
return ( result, )
|
||||
blip_model = BlipWrapper(caption_model_id=blip_model, vqa_model_id=vqa_model_id, device=device, cache_dir=blip_dir)
|
||||
|
||||
return ( blip_model, )
|
||||
|
||||
|
||||
# BLIP CAPTION IMAGE
|
||||
@@ -11031,105 +11014,47 @@ class WAS_BLIP_Analyze_Image:
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"image": ("IMAGE",),
|
||||
"images": ("IMAGE",),
|
||||
"mode": (["caption", "interrogate"], ),
|
||||
"question": ("STRING", {"default": "What does the background consist of?", "multiline": True}),
|
||||
"question": ("STRING", {"default": "What does the background consist of?", "multiline": True, "dynamicPrompts": False}),
|
||||
"blip_model": ("BLIP_MODEL",),
|
||||
},
|
||||
"optional": {
|
||||
"blip_model": ("BLIP_MODEL",)
|
||||
"min_length": ("INT", {"min": 1, "max": 1024, "default": 24}),
|
||||
"max_length": ("INT", {"min": 2, "max": 1024, "default": 64}),
|
||||
"num_beams": ("INT", {"min": 1, "max": 12, "default": 5}),
|
||||
"no_repeat_ngram_size": ("INT", {"min": 1, "max": 12, "default": 3}),
|
||||
"early_stopping": ("BOOLEAN", {"default": False})
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = (TEXT_TYPE,)
|
||||
FUNCTION = "blip_caption_image"
|
||||
RETURN_TYPES = (TEXT_TYPE, TEXT_TYPE)
|
||||
OUTPUT_IS_LIST = (False, True)
|
||||
|
||||
FUNCTION = "blip_caption_image"
|
||||
CATEGORY = "WAS Suite/Text/AI"
|
||||
|
||||
def blip_caption_image(self, image, mode, question, blip_model=None):
|
||||
def blip_caption_image(self, images, mode, question, blip_model, min_length=24, max_length=64, num_beams=5, no_repeat_ngram_size=3, early_stopping=False):
|
||||
|
||||
def transformImage(input_image, image_size, device):
|
||||
raw_image = input_image.convert('RGB')
|
||||
raw_image = raw_image.resize((image_size, image_size))
|
||||
transform = transforms.Compose([
|
||||
transforms.Resize(raw_image.size, interpolation=InterpolationMode.BICUBIC),
|
||||
transforms.ToTensor(),
|
||||
transforms.Normalize((0.48145466, 0.4578275, 0.40821073), (0.26862954, 0.26130258, 0.27577711))
|
||||
])
|
||||
image = transform(raw_image).unsqueeze(0).to(device)
|
||||
return image.view(1, -1, image_size, image_size) # Change the shape of the output tensor
|
||||
|
||||
from torchvision import transforms
|
||||
from torchvision.transforms.functional import InterpolationMode
|
||||
|
||||
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
|
||||
|
||||
conf = getSuiteConfig()
|
||||
image = tensor2pil(image)
|
||||
size = 384
|
||||
tensor = transformImage(image, size, device)
|
||||
|
||||
if blip_model:
|
||||
mode = blip_model[1]
|
||||
|
||||
if mode == 'caption':
|
||||
|
||||
if blip_model:
|
||||
model = blip_model[0].to(device)
|
||||
captions = []
|
||||
for image in images:
|
||||
|
||||
pil_image = tensor2pil(image).convert("RGB")
|
||||
|
||||
if mode == "caption":
|
||||
cap = blip_model.generate_caption(image=pil_image, min_length=min_length, max_length=max_length, num_beams=num_beams, no_repeat_ngram_size=no_repeat_ngram_size, early_stopping=early_stopping)
|
||||
captions.append(cap)
|
||||
cstr(f"\033[33mBLIP Caption:\033[0m {cap}").msg.print()
|
||||
else:
|
||||
from .modules.BLIP.blip_module import blip_decoder
|
||||
cap = blip_model.answer_question(image=pil_image, question=question, min_length=min_length, max_length=max_length, num_beams=num_beams, no_repeat_ngram_size=no_repeat_ngram_size, early_stopping=early_stopping)
|
||||
captions.append(cap)
|
||||
cstr(f"\033[33m BLIP Answer:\033[0m {cap}").msg.print()
|
||||
|
||||
blip_dir = os.path.join(MODELS_DIR, 'blip')
|
||||
if not os.path.exists(blip_dir):
|
||||
os.makedirs(blip_dir, exist_ok=True)
|
||||
full_captions = ""
|
||||
for i, caption in enumerate(captions):
|
||||
full_captions += caption + ("\n\n" if i < len(captions) else "")
|
||||
|
||||
torch.hub.set_dir(blip_dir)
|
||||
|
||||
if conf.__contains__('blip_model_url'):
|
||||
model_url = conf['blip_model_url']
|
||||
else:
|
||||
model_url = 'https://storage.googleapis.com/sfr-vision-language-research/BLIP/models/model_base_capfilt_large.pth'
|
||||
|
||||
model = blip_decoder(pretrained=model_url, image_size=size, vit='base')
|
||||
model.eval()
|
||||
model = model.to(device)
|
||||
|
||||
with torch.no_grad():
|
||||
caption = model.generate(tensor, sample=False, num_beams=6, max_length=74, min_length=20)
|
||||
# nucleus sampling
|
||||
#caption = model.generate(tensor, sample=True, top_p=0.9, max_length=75, min_length=10)
|
||||
cstr(f"\033[33mBLIP Caption:\033[0m {caption[0]}").msg.print()
|
||||
return (caption[0], )
|
||||
|
||||
elif mode == 'interrogate':
|
||||
|
||||
if blip_model:
|
||||
model = blip_model[0].to(device)
|
||||
else:
|
||||
from .modules.BLIP.blip_module import blip_vqa
|
||||
|
||||
blip_dir = os.path.join(MODELS_DIR, 'blip')
|
||||
if not os.path.exists(blip_dir):
|
||||
os.makedirs(blip_dir, exist_ok=True)
|
||||
|
||||
torch.hub.set_dir(blip_dir)
|
||||
|
||||
if conf.__contains__('blip_model_vqa_url'):
|
||||
model_url = conf['blip_model_vqa_url']
|
||||
else:
|
||||
model_url = 'https://storage.googleapis.com/sfr-vision-language-research/BLIP/models/model_base_vqa_capfilt_large.pth'
|
||||
|
||||
model = blip_vqa(pretrained=model_url, image_size=size, vit='base')
|
||||
model.eval()
|
||||
model = model.to(device)
|
||||
|
||||
with torch.no_grad():
|
||||
answer = model(tensor, question, train=False, inference='generate')
|
||||
cstr(f"\033[33m BLIP Answer:\033[0m {answer[0]}").msg.print()
|
||||
return (answer[0], )
|
||||
|
||||
else:
|
||||
cstr(f"The selected mode `{mode}` is not a valid selection!").error.print()
|
||||
return ('Invalid BLIP mode!', )
|
||||
return (full_captions, captions)
|
||||
|
||||
|
||||
# CLIPSeg Model Loader
|
||||
|
||||
+3
-3
@@ -7,14 +7,14 @@ imageio
|
||||
joblib
|
||||
matplotlib
|
||||
numba
|
||||
numpy>=1.18.5, <1.25.0
|
||||
numpy
|
||||
opencv-python-headless[ffmpeg]<=4.7.0.72
|
||||
pilgram
|
||||
git+https://github.com/WASasquatch/ffmpy.git
|
||||
rembg
|
||||
scikit-image==0.20.0
|
||||
scikit-image>=0.20.0
|
||||
scikit-learn
|
||||
scipy
|
||||
timm>=0.4.12
|
||||
tqdm
|
||||
transformers==4.26.1
|
||||
transformers
|
||||
|
||||
Reference in New Issue
Block a user