Skip to content

Commit 089d035

Browse files
committed
tmp save poc
1 parent 228607e commit 089d035

5 files changed

Lines changed: 47 additions & 37 deletions

File tree

bizyair_extras/nodes_flux_trainer.py

Lines changed: 0 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -1,8 +1,4 @@
11
import io
2-
import json
3-
import os
4-
import re
5-
import time
62

73
import matplotlib.pyplot as plt
84

nodes.py

Lines changed: 31 additions & 31 deletions
Original file line numberDiff line numberDiff line change
@@ -125,16 +125,16 @@ def INPUT_TYPES(s):
125125
}
126126

127127
RETURN_TYPES = ("LATENT",)
128-
FUNCTION = "sample"
128+
# FUNCTION = "sample"
129129

130130
CATEGORY = "sampling"
131131

132-
def sample(self, model, **kwargs):
133-
new_model: BizyAirNodeIO = model.copy(self.assigned_id)
134-
kwargs["model"] = model
135-
new_model.add_node_data(class_type="KSamplerAdvanced", inputs=kwargs)
136-
progress_callback = ProgressCallback()
137-
return new_model.send_request(progress_callback=progress_callback)
132+
# def sample(self, model, **kwargs):
133+
# new_model: BizyAirNodeIO = model.copy(self.assigned_id)
134+
# kwargs["model"] = model
135+
# new_model.add_node_data(class_type="KSamplerAdvanced", inputs=kwargs)
136+
# progress_callback = ProgressCallback()
137+
# return new_model.send_request(progress_callback=progress_callback)
138138

139139

140140
class BizyAir_CheckpointLoaderSimple(BizyAirBaseNode):
@@ -222,21 +222,21 @@ def INPUT_TYPES(s):
222222

223223
RETURN_TYPES = ("IMAGE",)
224224
RETURN_NAMES = (f"IMAGE",)
225-
FUNCTION = "decode"
225+
# FUNCTION = "decode"
226226

227227
CATEGORY = f"{PREFIX}/latent"
228228

229-
def decode(self, vae, samples):
230-
new_vae: BizyAirNodeIO = vae.copy(self.assigned_id)
231-
new_vae.add_node_data(
232-
class_type="VAEDecode",
233-
inputs={
234-
"samples": samples,
235-
"vae": vae,
236-
},
237-
outputs={"slot_index": 0},
238-
)
239-
return new_vae.send_request()
229+
# def decode(self, vae, samples):
230+
# new_vae: BizyAirNodeIO = vae.copy(self.assigned_id)
231+
# new_vae.add_node_data(
232+
# class_type="VAEDecode",
233+
# inputs={
234+
# "samples": samples,
235+
# "vae": vae,
236+
# },
237+
# outputs={"slot_index": 0},
238+
# )
239+
# return new_vae.send_request()
240240

241241

242242
class BizyAir_LoraLoader(BizyAirBaseNode):
@@ -295,20 +295,20 @@ def INPUT_TYPES(s):
295295

296296
RETURN_TYPES = ("LATENT",)
297297
RETURN_NAMES = (f"LATENT",)
298-
FUNCTION = "encode"
298+
# FUNCTION = "encode"
299299
CATEGORY = f"{PREFIX}/latent"
300300

301-
def encode(self, vae, pixels):
302-
new_vae: BizyAirNodeIO = vae.copy(self.assigned_id)
303-
new_vae.add_node_data(
304-
class_type="VAEEncode",
305-
inputs={
306-
"vae": vae,
307-
"pixels": pixels,
308-
},
309-
outputs={"slot_index": 0},
310-
)
311-
return new_vae.send_request()
301+
# def encode(self, vae, pixels):
302+
# new_vae: BizyAirNodeIO = vae.copy(self.assigned_id)
303+
# new_vae.add_node_data(
304+
# class_type="VAEEncode",
305+
# inputs={
306+
# "vae": vae,
307+
# "pixels": pixels,
308+
# },
309+
# outputs={"slot_index": 0},
310+
# )
311+
# return new_vae.send_request()
312312

313313

314314
class BizyAir_VAEEncodeForInpaint(BizyAirBaseNode):

src/bizyair/commands/servers/pub_sub_sse.py

Lines changed: 8 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -118,6 +118,7 @@ def __init__(self, name, mediator):
118118
self.mediator = mediator
119119
self.messages = queue.Queue()
120120
self.is_subscribed = True # 订阅状态监测
121+
self.tmp_result = {}
121122

122123
def receive(self, message):
123124
self.messages.put(message)
@@ -141,6 +142,11 @@ def unsubscribe(self):
141142

142143
def get_result(self, node_id, timeout=12):
143144
while True:
145+
if node_id in self.tmp_result:
146+
out = self.tmp_result[node_id]
147+
del self.tmp_result[node_id]
148+
return out
149+
144150
result = self.pop(timeout=timeout)
145151
if result is None:
146152
return None
@@ -153,6 +159,8 @@ def get_result(self, node_id, timeout=12):
153159
event_node_id = result["message"]["data"]["node"]
154160
if event_node_id == node_id:
155161
return result["data"]["payload"]
162+
else:
163+
self.tmp_result[event_node_id] = result["data"]["payload"]
156164
except Exception as e:
157165
print(f"Error processing message for {self.name}: {e}")
158166
return None

src/bizyair/image_utils.py

Lines changed: 6 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -176,14 +176,18 @@ def tensor_to_base64(tensor: torch.Tensor, compress=True) -> str:
176176
return tensor_b64
177177

178178

179-
def base64_to_tensor(tensor_b64: str, compress=True) -> torch.Tensor:
179+
def base64_to_numpy(tensor_b64: str, compress=True) -> np.ndarray:
180180
tensor_bytes = base64.b64decode(tensor_b64)
181181

182182
if compress:
183183
tensor_bytes = zlib.decompress(tensor_bytes)
184184

185185
tensor_np = pickle.loads(tensor_bytes)
186+
return tensor_np
187+
186188

189+
def base64_to_tensor(tensor_b64: str, compress=True) -> torch.Tensor:
190+
tensor_np = base64_to_numpy(tensor_b64, compress=compress)
187191
tensor = torch.from_numpy(tensor_np)
188192
return tensor
189193

@@ -222,7 +226,7 @@ def _(input: str, **kwargs):
222226
return decode_comfy_image(tensor_b64, old_version=old_version)
223227
elif input.startswith(NUMPY_MARKER):
224228
tensor_b64 = input[len(NUMPY_MARKER) :]
225-
return base64_to_tensor(tensor_b64)
229+
return base64_to_numpy(tensor_b64)
226230
return input
227231

228232

src/bizyair/nodes_base.py

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -166,6 +166,8 @@ def default_function(self, **kwargs):
166166
result = self._pre_run()
167167
if result:
168168
return self._merge_results(result, node_ios)
169+
else:
170+
warnings.warn("Pre-run result is None")
169171

170172
if len(send_request_datatype_list) == len(self.RETURN_TYPES):
171173
return self._process_all_send_request_types(node_ios)

0 commit comments

Comments
 (0)