diff --git a/__init__.py b/__init__.py index 15ccc33..6068c29 100644 --- a/__init__.py +++ b/__init__.py @@ -1,13 +1,19 @@ import importlib import importlib.util -from .pyserver import get_key_from_jssetting, update_models, update_styles +from .pyserver import ( + get_key_from_jssetting, + #update_characters, + update_models, + update_styles, +) node_list = [ "things_n_stuff_node", "gen_image_node", # "gen_image_inpaint_node", "gen_text_node", + "gen_text_advanced_node", "upscale_image_node", "util_nodes", ] diff --git a/js/nodespawn.js b/js/nodespawn.js index c698f3c..cd7e835 100644 --- a/js/nodespawn.js +++ b/js/nodespawn.js @@ -77,12 +77,13 @@ app.registerExtension({ }; } - if (nodeData.name === "GenerateText_VENICE") { + if (nodeData.name === "GenerateText_VENICE" || nodeData.name === "GenerateTextAdvanced_VENICE") { const originalOnNodeCreated = nodeType.prototype.onNodeCreated; nodeType.prototype.onNodeCreated = async function () { if (originalOnNodeCreated) { originalOnNodeCreated.apply(this); } + // Find the model widget const modelWidget = this.widgets.find(w => w.name === "model"); if (modelWidget) { @@ -114,6 +115,38 @@ app.registerExtension({ alert(`(VeniceAI.NodeSpawn) Failed to fetch text models:\n${error}`); } } + + // Find the character_slug widget // TODO: move this to a separate node depending on how character slugs are used (if model dependent eg) + // const characterSlugWidget = this.widgets.find(w => w.name === "model"); + // if (characterSlugWidget) { + // try { + // console.log("(VeniceAI.NodeSpawn) Trying to fetch character slugs..."); + // const response = await api.fetchApi("/veniceai/get_characters_list"); + + // if (!response.ok) { + // throw new Error(`HTTP error: ${response.status} ${response.statusText}`); + // } + + // const rawText = await response.text(); + + // let data; + // try { + // data = JSON.parse(rawText); + // } catch (jsonError) { + // throw new Error(`Failed to parse JSON: ${jsonError.message}. Raw response: ${rawText}`); + // } + + // characterSlugWidget.options.values = data.characters; // todo: implement on python side + // if (characterSlugWidget.onChange) { + // characterSlugWidget.onChange(); + // } + + // this.setDirtyCanvas(true); + // } catch (error) { + // console.error("(VeniceAI.NodeSpawn) Failed to fetch character slugs:", error); + // alert(`(VeniceAI.NodeSpawn) Failed to fetch character slugs:\n${error}`); + // } + // } }; } } diff --git a/js/startup.js b/js/startup.js index 544a052..10f0a9d 100644 --- a/js/startup.js +++ b/js/startup.js @@ -24,8 +24,8 @@ app.registerExtension({ alert(`${data.message}`); console.log(`(VeniceAI.Startup) ${data.message}`); } - else { - // update the model list if not model list error + else{ + // update the style list if not model list error console.log("(VeniceAI.Startup) Updating styles list..."); //alert("fetching styles list") const response_s = await api.fetchApi("/veniceai/update_styles_list"); @@ -35,6 +35,15 @@ app.registerExtension({ alert(`${data_s.message}`); console.log(`(VeniceAI.Startup) ${data_s.message}`); } + + // update the characters list + // console.log("(VeniceAI.Startup) Updating characters list..."); + // const response_c = await api.fetchApi("/veniceai/update_characters_list"); + // const data_c = await response_c.json(); + // if (data_c.error) { + // alert(`${data_c.message}`); + // console.log(`(VeniceAI.Startup) ${data_c.message}`); + // } } } catch (error) { // Handle any unexpected errors diff --git a/nodes/gen_text_advanced_node.py b/nodes/gen_text_advanced_node.py new file mode 100644 index 0000000..70e84e2 --- /dev/null +++ b/nodes/gen_text_advanced_node.py @@ -0,0 +1,293 @@ +import base64 +import io +import os + +import numpy as np +import requests +from PIL import Image + +from ..globals import API_ENDPOINTS, VENICEAI_BASE_URL + + +class GenerateTextAdvanced: + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "model": ("COMBO", {"default": "llama-3.3-70b"}), + "prompt": ("STRING", {"default": "", "multiline": True}), + "system_prompt": ("STRING", {"default": "", "multiline": True}), + "enable_system_prompt": ("BOOLEAN", {"default": True}), + # region venice_parameters + "enable_web_search": (["auto", "on", "off"], {"default": "auto"}), + "use_venice_system_prompt": ( + "BOOLEAN", + { + "default": False, + "tooltip": "Whether to include the Venice supplied system prompts along side specified system prompts.", + }, + ), + # "character_slug": ( + # "COMBO", + # {"default": "none", "tooltip": "The character slug of a public Venice character."}, + # ), + # endregion + "frequency_penalty": ( + "FLOAT", + { + "default": 0.0, + "min": -2.0, + "max": 2.0, + "step": 0.05, + "tooltip": "Positive values penalize new tokens based on their existing frequency in the text so far, decreasing the model's likelihood to repeat the same line verbatim.", + }, + ), + "presence_penalty": ( + "FLOAT", + { + "default": 0.0, + "min": -2.0, + "max": 2.0, + "step": 0.05, + "tooltip": "Positive values penalize new tokens based on whether they appear in the text so far, increasing the model's likelihood to talk about new topics.", + }, + ), + "repetition_penalty": ( + "FLOAT", + { + "default": 1.2, + "min": 0.0, + "max": 2.0, + "step": 0.05, + "tooltip": "1.0 means no penalty. Values > 1.0 discourage repetition.", + }, + ), + "max_temp": ( + "FLOAT", + { + "default": 1.5, + "min": 0.0, + "max": 2.0, + "step": 0.05, + "tooltip": "Maximum temperature value for dynamic temperature scaling.", + }, + ), + "min_temp": ( + "FLOAT", + { + "default": 0.1, + "min": 0.0, + "max": 2.0, + "step": 0.05, + "tooltip": "Minimum temperature value for dynamic temperature scaling.", + }, + ), + "max_completion_tokens": ( + "INT", + { + "default": 123, + "min": 1, + "max": 131072, + "step": 1, + "tooltip": "An upper bound for the number of tokens that can be generated for a completion, including visible output tokens and reasoning tokens.", + }, + ), + "temperature": ( + "FLOAT", + { + "default": 0.5, + "min": 0.0, + "max": 2.0, + "step": 0.05, + "tooltip": "Higher values like 0.8 will make the output more random, while lower values like 0.2 will make it more focused and deterministic. We generally recommend altering this or top_p but not both.", + }, + ), + "top_k": ( + "INT", + { + "default": 40, + "min": 0, + "tooltip": "The number of highest probability vocabulary tokens to keep for top-k-filtering.", + }, + ), + "top_p": ( + "FLOAT", + { + "default": 0.8, + "min": 0.0, + "max": 2.0, + "step": 0.01, + "tooltip": "An alternative to sampling with temperature, called nucleus sampling, where the model considers the results of the tokens with top_p probability mass. So 0.1 means only the tokens comprising the top 10% probability mass are considered.", + }, + ), + "min_p": ( + "FLOAT", + { + "default": 0.05, + "min": 0.0, + "max": 1.0, + "step": 0.01, + "tooltip": "Sets a minimum probability threshold for token selection. Tokens with probabilities below this value are filtered out.", + }, + ), + # "stop": ("STRING", {"default": "", "tooltip": "Up to 4 sequences where the API will stop generating further tokens. Defaults to null.", "placeholder": "stop: [\"\\n\"]"}), + # "stop_token_ids": ("STRING", {"default": "", "tooltip": "Array of token IDs where the API will stop generating further tokens. Example: [151643, 151645]", "placeholder": "151643, 151645, ..."}), + "enable_qwen25_vision": ("BOOLEAN", {"default": False}), + }, + "optional": { + "image_for_vision": ("IMAGE",), + }, + } + + RETURN_TYPES = ("STRING",) + RETURN_NAMES = ("response",) + FUNCTION = "generate_text" + CATEGORY = "venice.ai" + + def generate_text( + # region params + self, + model, + prompt, + system_prompt, + enable_system_prompt, + enable_web_search, + use_venice_system_prompt, + frequency_penalty, + presence_penalty, + repetition_penalty, + max_temp, + min_temp, + max_completion_tokens, + temperature, + top_k, + top_p, + min_p, + enable_qwen25_vision, + **kwargs, + # endregion + ): + + url = VENICEAI_BASE_URL + API_ENDPOINTS["text_generate"] + + user_content = [] + image_for_vision = kwargs.get("image_for_vision", None) + + if image_for_vision is not None and enable_qwen25_vision: + # Convert tensor to PIL Image + image_tensor = image_for_vision[0] # shape: (H, W, 3) + image_np = image_tensor.cpu().numpy() # Still in (H, W, 3) + image_np = (image_np * 255).astype(np.uint8) # Scale from [0, 1] to [0, 255] if needed + pil_image = Image.fromarray(image_np) + + # Resize image to meet constraints + original_width, original_height = pil_image.size + aspect_ratio = original_width / original_height + + # Determine target dimensions + if original_width > original_height: + target_width = 1024 + target_height = int(target_width / aspect_ratio) + if target_height < 256: + target_height = 256 + target_width = int(target_height * aspect_ratio) + else: + target_height = 1024 + target_width = int(target_height * aspect_ratio) + if target_width < 256: + target_width = 256 + target_height = int(target_width / aspect_ratio) + + # Round dimensions to multiples of 14 + def round_down_to_multiple(value, multiple): + return (value // multiple) * multiple + + target_width = round_down_to_multiple(target_width, 14) + target_height = round_down_to_multiple(target_height, 14) + + # Ensure minimum dimension is 256 after rounding + if min(target_width, target_height) < 256: + if target_width < target_height: + target_width = ((256 + 13) // 14) * 14 + target_height = round_down_to_multiple(int(target_width / aspect_ratio), 14) + else: + target_height = ((256 + 13) // 14) * 14 + target_width = round_down_to_multiple(int(target_height * aspect_ratio), 14) + + pil_image = pil_image.resize((target_width, target_height), Image.LANCZOS) + + # Convert to base64 and check size + buffered = io.BytesIO() + pil_image.save(buffered, format="PNG") + img_base64 = base64.b64encode(buffered.getvalue()).decode("utf-8") + + # Resize further if base64 exceeds 4.5MB + while len(img_base64) > 4500000: + scaling_factor = (4500000 / len(img_base64)) ** 0.5 + new_width = int(target_width * scaling_factor) + new_height = int(target_height * scaling_factor) + + new_width = max(round_down_to_multiple(new_width, 14), 256) + new_height = max(round_down_to_multiple(new_height, 14), 256) + + pil_image = pil_image.resize((new_width, new_height), Image.LANCZOS) + target_width, target_height = new_width, new_height + + buffered = io.BytesIO() + pil_image.save(buffered, format="PNG") + img_base64 = base64.b64encode(buffered.getvalue()).decode("utf-8") + + user_content.extend( + [ + {"type": "text", "text": prompt}, + {"type": "image_url", "image_url": {"url": f"data:image/png;base64,{img_base64}"}}, + ] + ) + else: + user_content.append({"type": "text", "text": prompt}) + + if not enable_system_prompt: + system_prompt = "" + + messages = [{"role": "system", "content": system_prompt}] + messages.append({"role": "user", "content": user_content}) + + payload = { + "model": model, + "messages": messages, + "venice_parameters": { + "enable_web_search": enable_web_search, + "include_venice_system_prompt": use_venice_system_prompt, + # "character_slug": "venice", + }, + "frequency_penalty": frequency_penalty, + "presence_penalty": presence_penalty, + "repetition_penalty": repetition_penalty, + "max_temp": max_temp, + "min_temp": min_temp, + "max_completion_tokens": max_completion_tokens, + "temperature": temperature, + "top_k": top_k, + "top_p": top_p, + "min_p": min_p, + } + + headers = {"Authorization": f"Bearer {os.getenv('VENICEAI_API_KEY')}", "Content-Type": "application/json"} + response = requests.post(url, json=payload, headers=headers) + + if response.status_code != 200: + raise requests.exceptions.HTTPError(f"HTTP error: {response.status_code}, Response: {response.text}") + + json_response = response.json() + content = json_response["choices"][0]["message"]["content"] + print(content) + return (content,) + + +NODE_CLASS_MAPPINGS = { + "GenerateTextAdvanced_VENICE": GenerateTextAdvanced, +} + +NODE_DISPLAY_NAME_MAPPINGS = { + "GenerateTextAdvanced_VENICE": "Generate Text Advanced BETA (Venice)", +} diff --git a/pyserver/update_characters.py b/pyserver/update_characters.py index 8900bbe..afe776c 100644 --- a/pyserver/update_characters.py +++ b/pyserver/update_characters.py @@ -15,6 +15,7 @@ data_dir = script_dir.parent / "data" data_dir.mkdir(exist_ok=True) characters_list_path = data_dir / "characters_list.json" +# TODO: unfinished async def fetch_characters_list(): try: