From 313bef0609b63af1780bcdd4051202733be9ef98 Mon Sep 17 00:00:00 2001 From: shadowcz007 Date: Sun, 7 Jan 2024 22:55:24 +0800 Subject: [PATCH] =?UTF-8?q?ClipInterrogator=E4=BC=98=E5=8C=96?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- nodes/ClipInterrogator.py | 78 ++++++++++++++-- web/javascript/clipInterrogator_mixlab.js | 104 ++++++++++++++++------ 2 files changed, 152 insertions(+), 30 deletions(-) diff --git a/nodes/ClipInterrogator.py b/nodes/ClipInterrogator.py index e0220bd..a86c6c9 100644 --- a/nodes/ClipInterrogator.py +++ b/nodes/ClipInterrogator.py @@ -6,6 +6,7 @@ import comfy.utils import numpy as np import json import torch +import random from transformers import AutoProcessor, BlipForConditionalGeneration @@ -58,7 +59,53 @@ def image_analysis_fn(ci,image): trending_ranks = {trending: sim for trending, sim in zip(top_trendings, ci.similarities(image_features, top_trendings))} flavor_ranks = {flavor: sim for flavor, sim in zip(top_flavors, ci.similarities(image_features, top_flavors))} - return medium_ranks, artist_ranks, movement_ranks, trending_ranks, flavor_ranks + return { + "medium_ranks":medium_ranks, + "artist_ranks":artist_ranks, + "movement_ranks":movement_ranks, + "trending_ranks":trending_ranks, + "flavor_ranks":flavor_ranks + } + + +def generate_sentences(data): + sentences = [] + + # Get the length of data + data_length = len(data) + + # Use a recursive function to handle variable-length data + def generate_recursive(index, current_sentence, current_score): + # Check if recursion is complete + if index == data_length: + sentences.append({"sentence": current_sentence, "score": current_score}) + return + + # Get the current level data + current_data = data[index] + + # Iterate through the current level data + for phrase in current_data: + sentence = current_sentence + ("," if current_sentence.strip() else "") + phrase + score = current_score + current_data[phrase] + generate_recursive(index + 1, sentence, score) + + # Start recursive generation of sentences + generate_recursive(0, "", 0) + + # Sort the generated sentences by score in descending order + sentences.sort(key=lambda x: x["score"], reverse=True) + + def get_random_elements(elements, num): + return random.sample(elements, num) + + ps = get_random_elements(sentences, 5) + ps = [s["sentence"] for s in sorted(ps, key=lambda x: x["score"], reverse=True)] + + return ps + + + def image_to_prompt(ci,image, mode): ci.config.chunk_size = 2048 if ci.config.clip_model_name == "ViT-L-14/openai" else 1024 @@ -86,17 +133,22 @@ class ClipInterrogator: "prompt_mode": (['fast','classic','best','negative'],), "image_analysis": (["off","on"],), }, + + # "optional":{ + # "output":("CLIPINTERROGATOR", {"multiline": True,"default": "", "dynamicPrompts": False}) + # }, + } - RETURN_TYPES = ("STRING",) - RETURN_NAMES = ("prompt",) + RETURN_TYPES = ("STRING","STRING",) + RETURN_NAMES = ("prompt","random_samples",) FUNCTION = "run" CATEGORY = "♾️Mixlab/prompt" OUTPUT_NODE = True INPUT_IS_LIST = True - OUTPUT_IS_LIST = (True,) + OUTPUT_IS_LIST = (True,True,) global ci ci = None def run(self,image,prompt_mode,image_analysis): @@ -156,5 +208,21 @@ class ClipInterrogator: ci.caption_offloaded = True # analysis_result=[] + # items = app.graph.getNodeById(31).widgets[2].value["items"] + + random_samples=[] - return {"ui":{"prompt": prompt_result,"analysis":analysis_result},"result": (prompt_result,)} \ No newline at end of file + for r in analysis_result: + random_sample = generate_sentences([r['medium_ranks'], r['artist_ranks'],r['movement_ranks'],r['trending_ranks'],r['flavor_ranks']]) + for s in random_sample: + random_samples.append(s) + # print(len(random_samples)) + # print('-----') + # print( random_samples) + return { + "ui":{ + "prompt": prompt_result, + "analysis":analysis_result, + "random_samples":random_samples + }, + "result": (prompt_result,random_samples,)} \ No newline at end of file diff --git a/web/javascript/clipInterrogator_mixlab.js b/web/javascript/clipInterrogator_mixlab.js index 29c4c12..973b08e 100644 --- a/web/javascript/clipInterrogator_mixlab.js +++ b/web/javascript/clipInterrogator_mixlab.js @@ -3,11 +3,57 @@ import { app } from '../../../scripts/app.js' import { ComfyWidgets } from '../../../scripts/widgets.js' import { $el } from '../../../scripts/ui.js' +function getRandomElements (arr, num) { + var result = [] + var len = arr.length + + for (var i = 0; i < num; i++) { + var randomIndex = Math.floor(Math.random() * len) + result.push(arr[randomIndex]) + } + + return result +} + +const createPrompt = (node, prompts, items, sample) => { + const w = ComfyWidgets['STRING']( + node, + 'text', + ['STRING', { multiline: true }], + app + ).widget + w.inputEl.readOnly = true + w.inputEl.style.opacity = 0.6 + + w.value = typeof prompts === 'string' ? prompts : prompts.join('\n\n') + + const w2 = ComfyWidgets['STRING']( + node, + 'text', + ['STRING', { multiline: true }], + app + ).widget + w2.inputEl.readOnly = true + w2.inputEl.style.opacity = 0.6 + + w2.value = typeof items === 'string' ? items : JSON.stringify(items, null, 2) + + const w3 = ComfyWidgets['STRING']( + node, + 'text', + ['STRING', { multiline: true }], + app + ).widget + w3.inputEl.readOnly = true + w3.inputEl.style.opacity = 0.6 + w3.value = typeof sample === 'string' ? sample : sample.join('\n\n') +} + app.registerExtension({ name: 'Mixlab.prompt.ClipInterrogator', async beforeRegisterNodeDef (nodeType, nodeData, app) { if (nodeData.name === 'ClipInterrogator') { - function populate (prompts, items) { + function populate (prompts, items, random_samples) { if (this.widgets) { for (let i = 0; i < this.widgets.length; i++) { if (this.widgets[i].type !== 'combo') this.widgets[i].onRemove?.() @@ -15,29 +61,9 @@ app.registerExtension({ this.widgets.length = 2 } - const w = ComfyWidgets['STRING']( - this, - 'text', - ['STRING', { multiline: true }], - app - ).widget - w.inputEl.readOnly = true - w.inputEl.style.opacity = 0.6 + createPrompt(this, prompts, items, random_samples) - w.value = prompts.join('\n\n') - - const w2 = ComfyWidgets['STRING']( - this, - 'text', - ['STRING', { multiline: true }], - app - ).widget - w2.inputEl.readOnly = true - w2.inputEl.style.opacity = 0.6 - - w2.value = JSON.stringify(items, null, 2) - - console.log('ClipInterrogator',w,w2) + // console.log('ClipInterrogator', w, w2) requestAnimationFrame(() => { const sz = this.computeSize() if (sz[0] < this.size[0]) { @@ -55,11 +81,39 @@ app.registerExtension({ const onExecuted = nodeType.prototype.onExecuted nodeType.prototype.onExecuted = function (message) { onExecuted?.apply(this, arguments) - console.log('##', message) - populate.call(this, message.prompt, message.analysis) + // console.log('##', message) + populate.call( + this, + message.prompt, + message.analysis, + message.random_samples + ) } this.serialize_widgets = true //需要保存参数 } + }, + async loadedGraphNode (node, app) { + // Fires every time a node is constructed + // You can modify widgets/add handlers/etc here + + if (node.type === 'ClipInterrogator') { + try { + + let widgets_values = node.widgets_values + console.log(widgets_values ) + try { + if (widgets_values[2] && widgets_values[3] && widgets_values[4]) + createPrompt( + node, + widgets_values[2], + widgets_values[3], + widgets_values[4] + ) + } catch (error) { + console.log(error) + } + } catch (error) {} + } } })