Refactor BLIP! Some some other bugs fixed as well.

This commit is contained in:
Jordan Thompson
2024-05-04 08:59:16 -07:00
parent 6629ddc410
commit cda4d6b8e2
2 changed files with 82 additions and 157 deletions
+79 -154
View File
@@ -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
View File
@@ -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