mirror of
https://git.datalinker.icu/comfyanonymous/ComfyUI
synced 2026-09-08 23:37:08 +08:00
Merge branch 'comfyanonymous:master' into master
This commit is contained in:
commit
da822b5057
@ -171,7 +171,10 @@ class Flux(nn.Module):
|
|||||||
pe = None
|
pe = None
|
||||||
|
|
||||||
blocks_replace = patches_replace.get("dit", {})
|
blocks_replace = patches_replace.get("dit", {})
|
||||||
|
transformer_options["total_blocks"] = len(self.double_blocks)
|
||||||
|
transformer_options["block_type"] = "double"
|
||||||
for i, block in enumerate(self.double_blocks):
|
for i, block in enumerate(self.double_blocks):
|
||||||
|
transformer_options["block_index"] = i
|
||||||
if ("double_block", i) in blocks_replace:
|
if ("double_block", i) in blocks_replace:
|
||||||
def block_wrap(args):
|
def block_wrap(args):
|
||||||
out = {}
|
out = {}
|
||||||
@ -215,7 +218,10 @@ class Flux(nn.Module):
|
|||||||
if self.params.global_modulation:
|
if self.params.global_modulation:
|
||||||
vec, _ = self.single_stream_modulation(vec_orig)
|
vec, _ = self.single_stream_modulation(vec_orig)
|
||||||
|
|
||||||
|
transformer_options["total_blocks"] = len(self.single_blocks)
|
||||||
|
transformer_options["block_type"] = "single"
|
||||||
for i, block in enumerate(self.single_blocks):
|
for i, block in enumerate(self.single_blocks):
|
||||||
|
transformer_options["block_index"] = i
|
||||||
if ("single_block", i) in blocks_replace:
|
if ("single_block", i) in blocks_replace:
|
||||||
def block_wrap(args):
|
def block_wrap(args):
|
||||||
out = {}
|
out = {}
|
||||||
|
|||||||
@ -509,7 +509,7 @@ class NextDiT(nn.Module):
|
|||||||
|
|
||||||
if self.pad_tokens_multiple is not None:
|
if self.pad_tokens_multiple is not None:
|
||||||
pad_extra = (-cap_feats.shape[1]) % self.pad_tokens_multiple
|
pad_extra = (-cap_feats.shape[1]) % self.pad_tokens_multiple
|
||||||
cap_feats = torch.cat((cap_feats, self.cap_pad_token.to(device=cap_feats.device, dtype=cap_feats.dtype).unsqueeze(0).repeat(cap_feats.shape[0], pad_extra, 1)), dim=1)
|
cap_feats = torch.cat((cap_feats, self.cap_pad_token.to(device=cap_feats.device, dtype=cap_feats.dtype, copy=True).unsqueeze(0).repeat(cap_feats.shape[0], pad_extra, 1)), dim=1)
|
||||||
|
|
||||||
cap_pos_ids = torch.zeros(bsz, cap_feats.shape[1], 3, dtype=torch.float32, device=device)
|
cap_pos_ids = torch.zeros(bsz, cap_feats.shape[1], 3, dtype=torch.float32, device=device)
|
||||||
cap_pos_ids[:, :, 0] = torch.arange(cap_feats.shape[1], dtype=torch.float32, device=device) + 1.0
|
cap_pos_ids[:, :, 0] = torch.arange(cap_feats.shape[1], dtype=torch.float32, device=device) + 1.0
|
||||||
@ -525,7 +525,7 @@ class NextDiT(nn.Module):
|
|||||||
|
|
||||||
if self.pad_tokens_multiple is not None:
|
if self.pad_tokens_multiple is not None:
|
||||||
pad_extra = (-x.shape[1]) % self.pad_tokens_multiple
|
pad_extra = (-x.shape[1]) % self.pad_tokens_multiple
|
||||||
x = torch.cat((x, self.x_pad_token.to(device=x.device, dtype=x.dtype).unsqueeze(0).repeat(x.shape[0], pad_extra, 1)), dim=1)
|
x = torch.cat((x, self.x_pad_token.to(device=x.device, dtype=x.dtype, copy=True).unsqueeze(0).repeat(x.shape[0], pad_extra, 1)), dim=1)
|
||||||
x_pos_ids = torch.nn.functional.pad(x_pos_ids, (0, 0, 0, pad_extra))
|
x_pos_ids = torch.nn.functional.pad(x_pos_ids, (0, 0, 0, pad_extra))
|
||||||
|
|
||||||
freqs_cis = self.rope_embedder(torch.cat((cap_pos_ids, x_pos_ids), dim=1)).movedim(1, 2)
|
freqs_cis = self.rope_embedder(torch.cat((cap_pos_ids, x_pos_ids), dim=1)).movedim(1, 2)
|
||||||
|
|||||||
@ -690,7 +690,7 @@ def load_models_gpu(models, memory_required=0, force_patch_weights=False, minimu
|
|||||||
loaded_memory = loaded_model.model_loaded_memory()
|
loaded_memory = loaded_model.model_loaded_memory()
|
||||||
current_free_mem = get_free_memory(torch_dev) + loaded_memory
|
current_free_mem = get_free_memory(torch_dev) + loaded_memory
|
||||||
|
|
||||||
lowvram_model_memory = max(128 * 1024 * 1024, (current_free_mem - minimum_memory_required), min(current_free_mem * MIN_WEIGHT_MEMORY_RATIO, current_free_mem - minimum_inference_memory()))
|
lowvram_model_memory = max(0, (current_free_mem - minimum_memory_required), min(current_free_mem * MIN_WEIGHT_MEMORY_RATIO, current_free_mem - minimum_inference_memory()))
|
||||||
lowvram_model_memory = lowvram_model_memory - loaded_memory
|
lowvram_model_memory = lowvram_model_memory - loaded_memory
|
||||||
|
|
||||||
if lowvram_model_memory == 0:
|
if lowvram_model_memory == 0:
|
||||||
@ -1013,7 +1013,7 @@ def force_channels_last():
|
|||||||
|
|
||||||
|
|
||||||
STREAMS = {}
|
STREAMS = {}
|
||||||
NUM_STREAMS = 1
|
NUM_STREAMS = 0
|
||||||
if args.async_offload:
|
if args.async_offload:
|
||||||
NUM_STREAMS = 2
|
NUM_STREAMS = 2
|
||||||
logging.info("Using async weight offloading with {} streams".format(NUM_STREAMS))
|
logging.info("Using async weight offloading with {} streams".format(NUM_STREAMS))
|
||||||
@ -1031,7 +1031,7 @@ def current_stream(device):
|
|||||||
stream_counters = {}
|
stream_counters = {}
|
||||||
def get_offload_stream(device):
|
def get_offload_stream(device):
|
||||||
stream_counter = stream_counters.get(device, 0)
|
stream_counter = stream_counters.get(device, 0)
|
||||||
if NUM_STREAMS <= 1:
|
if NUM_STREAMS == 0:
|
||||||
return None
|
return None
|
||||||
|
|
||||||
if device in STREAMS:
|
if device in STREAMS:
|
||||||
|
|||||||
@ -148,6 +148,15 @@ class LowVramPatch:
|
|||||||
else:
|
else:
|
||||||
return out
|
return out
|
||||||
|
|
||||||
|
#The above patch logic may cast up the weight to fp32, and do math. Go with fp32 x 3
|
||||||
|
LOWVRAM_PATCH_ESTIMATE_MATH_FACTOR = 3
|
||||||
|
|
||||||
|
def low_vram_patch_estimate_vram(model, key):
|
||||||
|
weight, set_func, convert_func = get_key_weight(model, key)
|
||||||
|
if weight is None:
|
||||||
|
return 0
|
||||||
|
return weight.numel() * torch.float32.itemsize * LOWVRAM_PATCH_ESTIMATE_MATH_FACTOR
|
||||||
|
|
||||||
def get_key_weight(model, key):
|
def get_key_weight(model, key):
|
||||||
set_func = None
|
set_func = None
|
||||||
convert_func = None
|
convert_func = None
|
||||||
@ -269,6 +278,9 @@ class ModelPatcher:
|
|||||||
if not hasattr(self.model, 'current_weight_patches_uuid'):
|
if not hasattr(self.model, 'current_weight_patches_uuid'):
|
||||||
self.model.current_weight_patches_uuid = None
|
self.model.current_weight_patches_uuid = None
|
||||||
|
|
||||||
|
if not hasattr(self.model, 'model_offload_buffer_memory'):
|
||||||
|
self.model.model_offload_buffer_memory = 0
|
||||||
|
|
||||||
def model_size(self):
|
def model_size(self):
|
||||||
if self.size > 0:
|
if self.size > 0:
|
||||||
return self.size
|
return self.size
|
||||||
@ -662,7 +674,16 @@ class ModelPatcher:
|
|||||||
skip = True # skip random weights in non leaf modules
|
skip = True # skip random weights in non leaf modules
|
||||||
break
|
break
|
||||||
if not skip and (hasattr(m, "comfy_cast_weights") or len(params) > 0):
|
if not skip and (hasattr(m, "comfy_cast_weights") or len(params) > 0):
|
||||||
loading.append((comfy.model_management.module_size(m), n, m, params))
|
module_mem = comfy.model_management.module_size(m)
|
||||||
|
module_offload_mem = module_mem
|
||||||
|
if hasattr(m, "comfy_cast_weights"):
|
||||||
|
weight_key = "{}.weight".format(n)
|
||||||
|
bias_key = "{}.bias".format(n)
|
||||||
|
if weight_key in self.patches:
|
||||||
|
module_offload_mem += low_vram_patch_estimate_vram(self.model, weight_key)
|
||||||
|
if bias_key in self.patches:
|
||||||
|
module_offload_mem += low_vram_patch_estimate_vram(self.model, bias_key)
|
||||||
|
loading.append((module_offload_mem, module_mem, n, m, params))
|
||||||
return loading
|
return loading
|
||||||
|
|
||||||
def load(self, device_to=None, lowvram_model_memory=0, force_patch_weights=False, full_load=False):
|
def load(self, device_to=None, lowvram_model_memory=0, force_patch_weights=False, full_load=False):
|
||||||
@ -676,20 +697,22 @@ class ModelPatcher:
|
|||||||
|
|
||||||
load_completely = []
|
load_completely = []
|
||||||
offloaded = []
|
offloaded = []
|
||||||
|
offload_buffer = 0
|
||||||
loading.sort(reverse=True)
|
loading.sort(reverse=True)
|
||||||
for x in loading:
|
for x in loading:
|
||||||
n = x[1]
|
module_offload_mem, module_mem, n, m, params = x
|
||||||
m = x[2]
|
|
||||||
params = x[3]
|
|
||||||
module_mem = x[0]
|
|
||||||
|
|
||||||
lowvram_weight = False
|
lowvram_weight = False
|
||||||
|
|
||||||
|
potential_offload = max(offload_buffer, module_offload_mem * (comfy.model_management.NUM_STREAMS + 1))
|
||||||
|
lowvram_fits = mem_counter + module_mem + potential_offload < lowvram_model_memory
|
||||||
|
|
||||||
weight_key = "{}.weight".format(n)
|
weight_key = "{}.weight".format(n)
|
||||||
bias_key = "{}.bias".format(n)
|
bias_key = "{}.bias".format(n)
|
||||||
|
|
||||||
if not full_load and hasattr(m, "comfy_cast_weights"):
|
if not full_load and hasattr(m, "comfy_cast_weights"):
|
||||||
if mem_counter + module_mem >= lowvram_model_memory:
|
if not lowvram_fits:
|
||||||
|
offload_buffer = potential_offload
|
||||||
lowvram_weight = True
|
lowvram_weight = True
|
||||||
lowvram_counter += 1
|
lowvram_counter += 1
|
||||||
lowvram_mem_counter += module_mem
|
lowvram_mem_counter += module_mem
|
||||||
@ -723,9 +746,11 @@ class ModelPatcher:
|
|||||||
if hasattr(m, "comfy_cast_weights"):
|
if hasattr(m, "comfy_cast_weights"):
|
||||||
wipe_lowvram_weight(m)
|
wipe_lowvram_weight(m)
|
||||||
|
|
||||||
if full_load or mem_counter + module_mem < lowvram_model_memory:
|
if full_load or lowvram_fits:
|
||||||
mem_counter += module_mem
|
mem_counter += module_mem
|
||||||
load_completely.append((module_mem, n, m, params))
|
load_completely.append((module_mem, n, m, params))
|
||||||
|
else:
|
||||||
|
offload_buffer = potential_offload
|
||||||
|
|
||||||
if cast_weight and hasattr(m, "comfy_cast_weights"):
|
if cast_weight and hasattr(m, "comfy_cast_weights"):
|
||||||
m.prev_comfy_cast_weights = m.comfy_cast_weights
|
m.prev_comfy_cast_weights = m.comfy_cast_weights
|
||||||
@ -766,7 +791,7 @@ class ModelPatcher:
|
|||||||
self.pin_weight_to_device("{}.{}".format(n, param))
|
self.pin_weight_to_device("{}.{}".format(n, param))
|
||||||
|
|
||||||
if lowvram_counter > 0:
|
if lowvram_counter > 0:
|
||||||
logging.info("loaded partially; {:.2f} MB usable, {:.2f} MB loaded, {:.2f} MB offloaded, lowvram patches: {}".format(lowvram_model_memory / (1024 * 1024), mem_counter / (1024 * 1024), lowvram_mem_counter / (1024 * 1024), patch_counter))
|
logging.info("loaded partially; {:.2f} MB usable, {:.2f} MB loaded, {:.2f} MB offloaded, {:.2f} MB buffer reserved, lowvram patches: {}".format(lowvram_model_memory / (1024 * 1024), mem_counter / (1024 * 1024), lowvram_mem_counter / (1024 * 1024), offload_buffer / (1024 * 1024), patch_counter))
|
||||||
self.model.model_lowvram = True
|
self.model.model_lowvram = True
|
||||||
else:
|
else:
|
||||||
logging.info("loaded completely; {:.2f} MB usable, {:.2f} MB loaded, full load: {}".format(lowvram_model_memory / (1024 * 1024), mem_counter / (1024 * 1024), full_load))
|
logging.info("loaded completely; {:.2f} MB usable, {:.2f} MB loaded, full load: {}".format(lowvram_model_memory / (1024 * 1024), mem_counter / (1024 * 1024), full_load))
|
||||||
@ -778,6 +803,7 @@ class ModelPatcher:
|
|||||||
self.model.lowvram_patch_counter += patch_counter
|
self.model.lowvram_patch_counter += patch_counter
|
||||||
self.model.device = device_to
|
self.model.device = device_to
|
||||||
self.model.model_loaded_weight_memory = mem_counter
|
self.model.model_loaded_weight_memory = mem_counter
|
||||||
|
self.model.model_offload_buffer_memory = offload_buffer
|
||||||
self.model.current_weight_patches_uuid = self.patches_uuid
|
self.model.current_weight_patches_uuid = self.patches_uuid
|
||||||
|
|
||||||
for callback in self.get_all_callbacks(CallbacksMP.ON_LOAD):
|
for callback in self.get_all_callbacks(CallbacksMP.ON_LOAD):
|
||||||
@ -831,6 +857,7 @@ class ModelPatcher:
|
|||||||
self.model.to(device_to)
|
self.model.to(device_to)
|
||||||
self.model.device = device_to
|
self.model.device = device_to
|
||||||
self.model.model_loaded_weight_memory = 0
|
self.model.model_loaded_weight_memory = 0
|
||||||
|
self.model.model_offload_buffer_memory = 0
|
||||||
|
|
||||||
for m in self.model.modules():
|
for m in self.model.modules():
|
||||||
if hasattr(m, "comfy_patched_weights"):
|
if hasattr(m, "comfy_patched_weights"):
|
||||||
@ -849,13 +876,14 @@ class ModelPatcher:
|
|||||||
patch_counter = 0
|
patch_counter = 0
|
||||||
unload_list = self._load_list()
|
unload_list = self._load_list()
|
||||||
unload_list.sort()
|
unload_list.sort()
|
||||||
|
offload_buffer = self.model.model_offload_buffer_memory
|
||||||
|
|
||||||
for unload in unload_list:
|
for unload in unload_list:
|
||||||
if memory_to_free < memory_freed:
|
if memory_to_free + offload_buffer - self.model.model_offload_buffer_memory < memory_freed:
|
||||||
break
|
break
|
||||||
module_mem = unload[0]
|
module_offload_mem, module_mem, n, m, params = unload
|
||||||
n = unload[1]
|
|
||||||
m = unload[2]
|
potential_offload = (comfy.model_management.NUM_STREAMS + 1) * module_offload_mem
|
||||||
params = unload[3]
|
|
||||||
|
|
||||||
lowvram_possible = hasattr(m, "comfy_cast_weights")
|
lowvram_possible = hasattr(m, "comfy_cast_weights")
|
||||||
if hasattr(m, "comfy_patched_weights") and m.comfy_patched_weights == True:
|
if hasattr(m, "comfy_patched_weights") and m.comfy_patched_weights == True:
|
||||||
@ -906,15 +934,18 @@ class ModelPatcher:
|
|||||||
m.comfy_cast_weights = True
|
m.comfy_cast_weights = True
|
||||||
m.comfy_patched_weights = False
|
m.comfy_patched_weights = False
|
||||||
memory_freed += module_mem
|
memory_freed += module_mem
|
||||||
|
offload_buffer = max(offload_buffer, potential_offload)
|
||||||
logging.debug("freed {}".format(n))
|
logging.debug("freed {}".format(n))
|
||||||
|
|
||||||
for param in params:
|
for param in params:
|
||||||
self.pin_weight_to_device("{}.{}".format(n, param))
|
self.pin_weight_to_device("{}.{}".format(n, param))
|
||||||
|
|
||||||
|
|
||||||
self.model.model_lowvram = True
|
self.model.model_lowvram = True
|
||||||
self.model.lowvram_patch_counter += patch_counter
|
self.model.lowvram_patch_counter += patch_counter
|
||||||
self.model.model_loaded_weight_memory -= memory_freed
|
self.model.model_loaded_weight_memory -= memory_freed
|
||||||
logging.info("loaded partially: {:.2f} MB loaded, lowvram patches: {}".format(self.model.model_loaded_weight_memory / (1024 * 1024), self.model.lowvram_patch_counter))
|
self.model.model_offload_buffer_memory = offload_buffer
|
||||||
|
logging.info("Unloaded partially: {:.2f} MB freed, {:.2f} MB remains loaded, {:.2f} MB buffer reserved, lowvram patches: {}".format(memory_freed / (1024 * 1024), self.model.model_loaded_weight_memory / (1024 * 1024), offload_buffer / (1024 * 1024), self.model.lowvram_patch_counter))
|
||||||
return memory_freed
|
return memory_freed
|
||||||
|
|
||||||
def partially_load(self, device_to, extra_memory=0, force_patch_weights=False):
|
def partially_load(self, device_to, extra_memory=0, force_patch_weights=False):
|
||||||
|
|||||||
@ -4,10 +4,7 @@ See: https://cloud.google.com/vertex-ai/generative-ai/docs/model-reference/infer
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
import base64
|
import base64
|
||||||
import json
|
|
||||||
import os
|
import os
|
||||||
import time
|
|
||||||
import uuid
|
|
||||||
from enum import Enum
|
from enum import Enum
|
||||||
from io import BytesIO
|
from io import BytesIO
|
||||||
from typing import Literal
|
from typing import Literal
|
||||||
@ -43,7 +40,6 @@ from comfy_api_nodes.util import (
|
|||||||
validate_string,
|
validate_string,
|
||||||
video_to_base64_string,
|
video_to_base64_string,
|
||||||
)
|
)
|
||||||
from server import PromptServer
|
|
||||||
|
|
||||||
GEMINI_BASE_ENDPOINT = "/proxy/vertexai/gemini"
|
GEMINI_BASE_ENDPOINT = "/proxy/vertexai/gemini"
|
||||||
GEMINI_MAX_INPUT_FILE_SIZE = 20 * 1024 * 1024 # 20 MB
|
GEMINI_MAX_INPUT_FILE_SIZE = 20 * 1024 * 1024 # 20 MB
|
||||||
@ -384,29 +380,6 @@ class GeminiNode(IO.ComfyNode):
|
|||||||
)
|
)
|
||||||
|
|
||||||
output_text = get_text_from_response(response)
|
output_text = get_text_from_response(response)
|
||||||
if output_text:
|
|
||||||
# Not a true chat history like the OpenAI Chat node. It is emulated so the frontend can show a copy button.
|
|
||||||
render_spec = {
|
|
||||||
"node_id": cls.hidden.unique_id,
|
|
||||||
"component": "ChatHistoryWidget",
|
|
||||||
"props": {
|
|
||||||
"history": json.dumps(
|
|
||||||
[
|
|
||||||
{
|
|
||||||
"prompt": prompt,
|
|
||||||
"response": output_text,
|
|
||||||
"response_id": str(uuid.uuid4()),
|
|
||||||
"timestamp": time.time(),
|
|
||||||
}
|
|
||||||
]
|
|
||||||
),
|
|
||||||
},
|
|
||||||
}
|
|
||||||
PromptServer.instance.send_sync(
|
|
||||||
"display_component",
|
|
||||||
render_spec,
|
|
||||||
)
|
|
||||||
|
|
||||||
return IO.NodeOutput(output_text or "Empty response from Gemini model...")
|
return IO.NodeOutput(output_text or "Empty response from Gemini model...")
|
||||||
|
|
||||||
|
|
||||||
@ -601,30 +574,7 @@ class GeminiImage(IO.ComfyNode):
|
|||||||
response_model=GeminiGenerateContentResponse,
|
response_model=GeminiGenerateContentResponse,
|
||||||
price_extractor=calculate_tokens_price,
|
price_extractor=calculate_tokens_price,
|
||||||
)
|
)
|
||||||
|
return IO.NodeOutput(get_image_from_response(response), get_text_from_response(response))
|
||||||
output_text = get_text_from_response(response)
|
|
||||||
if output_text:
|
|
||||||
render_spec = {
|
|
||||||
"node_id": cls.hidden.unique_id,
|
|
||||||
"component": "ChatHistoryWidget",
|
|
||||||
"props": {
|
|
||||||
"history": json.dumps(
|
|
||||||
[
|
|
||||||
{
|
|
||||||
"prompt": prompt,
|
|
||||||
"response": output_text,
|
|
||||||
"response_id": str(uuid.uuid4()),
|
|
||||||
"timestamp": time.time(),
|
|
||||||
}
|
|
||||||
]
|
|
||||||
),
|
|
||||||
},
|
|
||||||
}
|
|
||||||
PromptServer.instance.send_sync(
|
|
||||||
"display_component",
|
|
||||||
render_spec,
|
|
||||||
)
|
|
||||||
return IO.NodeOutput(get_image_from_response(response), output_text)
|
|
||||||
|
|
||||||
|
|
||||||
class GeminiImage2(IO.ComfyNode):
|
class GeminiImage2(IO.ComfyNode):
|
||||||
@ -744,30 +694,7 @@ class GeminiImage2(IO.ComfyNode):
|
|||||||
response_model=GeminiGenerateContentResponse,
|
response_model=GeminiGenerateContentResponse,
|
||||||
price_extractor=calculate_tokens_price,
|
price_extractor=calculate_tokens_price,
|
||||||
)
|
)
|
||||||
|
return IO.NodeOutput(get_image_from_response(response), get_text_from_response(response))
|
||||||
output_text = get_text_from_response(response)
|
|
||||||
if output_text:
|
|
||||||
render_spec = {
|
|
||||||
"node_id": cls.hidden.unique_id,
|
|
||||||
"component": "ChatHistoryWidget",
|
|
||||||
"props": {
|
|
||||||
"history": json.dumps(
|
|
||||||
[
|
|
||||||
{
|
|
||||||
"prompt": prompt,
|
|
||||||
"response": output_text,
|
|
||||||
"response_id": str(uuid.uuid4()),
|
|
||||||
"timestamp": time.time(),
|
|
||||||
}
|
|
||||||
]
|
|
||||||
),
|
|
||||||
},
|
|
||||||
}
|
|
||||||
PromptServer.instance.send_sync(
|
|
||||||
"display_component",
|
|
||||||
render_spec,
|
|
||||||
)
|
|
||||||
return IO.NodeOutput(get_image_from_response(response), output_text)
|
|
||||||
|
|
||||||
|
|
||||||
class GeminiExtension(ComfyExtension):
|
class GeminiExtension(ComfyExtension):
|
||||||
|
|||||||
@ -1,15 +1,10 @@
|
|||||||
from io import BytesIO
|
from io import BytesIO
|
||||||
from typing import Optional, Union
|
|
||||||
import json
|
|
||||||
import os
|
import os
|
||||||
import time
|
|
||||||
import uuid
|
|
||||||
from enum import Enum
|
from enum import Enum
|
||||||
from inspect import cleandoc
|
from inspect import cleandoc
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import torch
|
import torch
|
||||||
from PIL import Image
|
from PIL import Image
|
||||||
from server import PromptServer
|
|
||||||
import folder_paths
|
import folder_paths
|
||||||
import base64
|
import base64
|
||||||
from comfy_api.latest import IO, ComfyExtension
|
from comfy_api.latest import IO, ComfyExtension
|
||||||
@ -587,11 +582,11 @@ class OpenAIChatNode(IO.ComfyNode):
|
|||||||
def create_input_message_contents(
|
def create_input_message_contents(
|
||||||
cls,
|
cls,
|
||||||
prompt: str,
|
prompt: str,
|
||||||
image: Optional[torch.Tensor] = None,
|
image: torch.Tensor | None = None,
|
||||||
files: Optional[list[InputFileContent]] = None,
|
files: list[InputFileContent] | None = None,
|
||||||
) -> InputMessageContentList:
|
) -> InputMessageContentList:
|
||||||
"""Create a list of input message contents from prompt and optional image."""
|
"""Create a list of input message contents from prompt and optional image."""
|
||||||
content_list: list[Union[InputContent, InputTextContent, InputImageContent, InputFileContent]] = [
|
content_list: list[InputContent | InputTextContent | InputImageContent | InputFileContent] = [
|
||||||
InputTextContent(text=prompt, type="input_text"),
|
InputTextContent(text=prompt, type="input_text"),
|
||||||
]
|
]
|
||||||
if image is not None:
|
if image is not None:
|
||||||
@ -617,9 +612,9 @@ class OpenAIChatNode(IO.ComfyNode):
|
|||||||
prompt: str,
|
prompt: str,
|
||||||
persist_context: bool = False,
|
persist_context: bool = False,
|
||||||
model: SupportedOpenAIModel = SupportedOpenAIModel.gpt_5.value,
|
model: SupportedOpenAIModel = SupportedOpenAIModel.gpt_5.value,
|
||||||
images: Optional[torch.Tensor] = None,
|
images: torch.Tensor | None = None,
|
||||||
files: Optional[list[InputFileContent]] = None,
|
files: list[InputFileContent] | None = None,
|
||||||
advanced_options: Optional[CreateModelResponseProperties] = None,
|
advanced_options: CreateModelResponseProperties | None = None,
|
||||||
) -> IO.NodeOutput:
|
) -> IO.NodeOutput:
|
||||||
validate_string(prompt, strip_whitespace=False)
|
validate_string(prompt, strip_whitespace=False)
|
||||||
|
|
||||||
@ -660,30 +655,7 @@ class OpenAIChatNode(IO.ComfyNode):
|
|||||||
status_extractor=lambda response: response.status,
|
status_extractor=lambda response: response.status,
|
||||||
completed_statuses=["incomplete", "completed"]
|
completed_statuses=["incomplete", "completed"]
|
||||||
)
|
)
|
||||||
output_text = cls.get_text_from_message_content(cls.get_message_content_from_response(result_response))
|
return IO.NodeOutput(cls.get_text_from_message_content(cls.get_message_content_from_response(result_response)))
|
||||||
|
|
||||||
# Update history
|
|
||||||
render_spec = {
|
|
||||||
"node_id": cls.hidden.unique_id,
|
|
||||||
"component": "ChatHistoryWidget",
|
|
||||||
"props": {
|
|
||||||
"history": json.dumps(
|
|
||||||
[
|
|
||||||
{
|
|
||||||
"prompt": prompt,
|
|
||||||
"response": output_text,
|
|
||||||
"response_id": str(uuid.uuid4()),
|
|
||||||
"timestamp": time.time(),
|
|
||||||
}
|
|
||||||
]
|
|
||||||
),
|
|
||||||
},
|
|
||||||
}
|
|
||||||
PromptServer.instance.send_sync(
|
|
||||||
"display_component",
|
|
||||||
render_spec,
|
|
||||||
)
|
|
||||||
return IO.NodeOutput(output_text)
|
|
||||||
|
|
||||||
|
|
||||||
class OpenAIInputFiles(IO.ComfyNode):
|
class OpenAIInputFiles(IO.ComfyNode):
|
||||||
@ -790,8 +762,8 @@ class OpenAIChatConfig(IO.ComfyNode):
|
|||||||
def execute(
|
def execute(
|
||||||
cls,
|
cls,
|
||||||
truncation: bool,
|
truncation: bool,
|
||||||
instructions: Optional[str] = None,
|
instructions: str | None = None,
|
||||||
max_output_tokens: Optional[int] = None,
|
max_output_tokens: int | None = None,
|
||||||
) -> IO.NodeOutput:
|
) -> IO.NodeOutput:
|
||||||
"""
|
"""
|
||||||
Configure advanced options for the OpenAI Chat Node.
|
Configure advanced options for the OpenAI Chat Node.
|
||||||
|
|||||||
File diff suppressed because it is too large
Load Diff
1432
comfy_extras/nodes_dataset.py
Normal file
1432
comfy_extras/nodes_dataset.py
Normal file
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
1
nodes.py
1
nodes.py
@ -2278,6 +2278,7 @@ async def init_builtin_extra_nodes():
|
|||||||
"nodes_images.py",
|
"nodes_images.py",
|
||||||
"nodes_video_model.py",
|
"nodes_video_model.py",
|
||||||
"nodes_train.py",
|
"nodes_train.py",
|
||||||
|
"nodes_dataset.py",
|
||||||
"nodes_sag.py",
|
"nodes_sag.py",
|
||||||
"nodes_perpneg.py",
|
"nodes_perpneg.py",
|
||||||
"nodes_stable3d.py",
|
"nodes_stable3d.py",
|
||||||
|
|||||||
Loading…
x
Reference in New Issue
Block a user