Refactoring and cleanups
This commit is contained in:
@@ -18,7 +18,7 @@ Restart ComfyUI to apply new changes.
|
||||
* Supports setting max preview size (ComfyUI default is hardcoded to 512 max).
|
||||
* Supports showing previews for more than the first latent in the batch.
|
||||
* Supports throttling previews. Do you really need your expensive TAESD preview to get updated 3 times a second?
|
||||
* Supports using CUDA streams to avoid waiting for a synchronize. This might be faster.
|
||||
* Supports using CUDA streams to avoid waiting for a synchronize. Increases speed slightly at the cost of higher VRAM usage. For comparison, a batch of 8 768x768 images with throttle at `0.5` sec is `1.29s/it` with it on and `1.42s/it` with it off for me.
|
||||
|
||||
Current defaults from `blehconfig.json`
|
||||
|
||||
@@ -29,9 +29,9 @@ Current defaults from `blehconfig.json`
|
||||
|`max_batch`|`4`|Max number of latents in a batch to preview|
|
||||
|`max_batch_cols`|`2`|Max number of columns to use when previewing batches|
|
||||
|`throttle_secs`|`1`|Max frequency to decode the latents for previewing. `0.25` would be every 1/4 sec, `2` would be only once every two seconds|
|
||||
|`use_cuda`|`true`|Use special logic for CUDA (and maybe pretend-CUDA like ROCM) to reduce the performance impact of preview generation|
|
||||
|`use_cuda`|`false`|Use special logic for CUDA (and maybe pretend-CUDA like ROCM) to reduce the performance impact of preview generation|
|
||||
|
||||
I would recommend setting `throttle_secs` to something relatively high like 5-10 sec especially if you are generating batches at high resolution.
|
||||
I would recommend setting `throttle_secs` to something relatively high like 5-10 sec especially if you are generating batches at high resolution. `use_cuda` now defaults to `false` as it can substantially increase VRAM requirements.
|
||||
|
||||
### BlehHyperTile
|
||||
|
||||
|
||||
+4
-3
@@ -1,17 +1,18 @@
|
||||
from .py import settings
|
||||
|
||||
settings.load_settings()
|
||||
|
||||
if settings.SETTINGS.btp_enabled:
|
||||
from .py import betterTaesdPreview
|
||||
from .py import betterTaesdPreview
|
||||
|
||||
from .py import hypertile
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"BlehHyperTile": hypertile.HyperTileBleh,
|
||||
"BlehHyperTile": hypertile.HyperTileBleh,
|
||||
}
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"HyperTile (bleh)": "HyperTile (bleh)",
|
||||
"HyperTile (bleh)": "HyperTile (bleh)",
|
||||
}
|
||||
|
||||
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
|
||||
|
||||
+91
-65
@@ -1,85 +1,111 @@
|
||||
import math
|
||||
from time import time
|
||||
|
||||
import latent_preview
|
||||
import numpy as np
|
||||
import torch
|
||||
from PIL import Image
|
||||
|
||||
import latent_preview
|
||||
|
||||
from .settings import SETTINGS
|
||||
|
||||
_ORIG_PREVIEWER = latent_preview.TAESDPreviewerImpl
|
||||
|
||||
|
||||
class BetterTAESDPreviewer(_ORIG_PREVIEWER):
|
||||
def __init__(self, taesd):
|
||||
self.taesd = taesd
|
||||
self.stamp = None
|
||||
self.cached = None
|
||||
self.blank = Image.new("RGB", size=(1,1))
|
||||
self.stream = None
|
||||
self.prev_work = None
|
||||
self.use_cuda = SETTINGS.btp_use_cuda and hasattr(torch, "cuda") and torch.cuda.is_available()
|
||||
def __init__(self, taesd):
|
||||
self.taesd = taesd
|
||||
self.stamp = None
|
||||
self.cached = None
|
||||
self.blank = Image.new("RGB", size=(1, 1))
|
||||
self.stream = None
|
||||
self.prev_work = None
|
||||
self.cpudev = torch.device("cpu")
|
||||
self.use_cuda = (
|
||||
SETTINGS.btp_use_cuda
|
||||
and hasattr(torch, "cuda")
|
||||
and torch.cuda.is_available()
|
||||
)
|
||||
|
||||
def decode_latent_to_preview_image(self, preview_format, x0):
|
||||
preview_image = self.decode_latent_to_preview(x0)
|
||||
return (preview_format, preview_image, min(max(*preview_image.size), SETTINGS.btp_max_size))
|
||||
def decode_latent_to_preview_image(self, preview_format, x0):
|
||||
preview_image = self.decode_latent_to_preview(x0)
|
||||
return (
|
||||
preview_format,
|
||||
preview_image,
|
||||
min(max(*preview_image.size), SETTINGS.btp_max_size),
|
||||
)
|
||||
|
||||
def check_use_cached(self):
|
||||
now = time()
|
||||
if self.cached is not None and self.stamp is not None:
|
||||
if now - self.stamp < SETTINGS.btp_throttle_secs:
|
||||
return True
|
||||
self.stamp = now
|
||||
return False
|
||||
def check_use_cached(self):
|
||||
now = time()
|
||||
if (
|
||||
self.cached is not None and self.stamp is not None
|
||||
) and now - self.stamp < SETTINGS.btp_throttle_secs:
|
||||
return True
|
||||
self.stamp = now
|
||||
return False
|
||||
|
||||
def _decode_latent(self, x0):
|
||||
samples = (self.taesd.decode(x0[:SETTINGS.btp_max_batch]) + 1.0) / 2.0
|
||||
samples = torch.clamp(samples, min = 0.0, max = 1.0) * 255.0
|
||||
return samples.to(dtype = torch.uint8).detach()
|
||||
def _decode_latent(self, x0):
|
||||
samples = (self.taesd.decode(x0[: SETTINGS.btp_max_batch]) + 1.0) / 2.0
|
||||
samples = torch.clamp(samples, min=0.0, max=1.0) * 255.0
|
||||
return samples.to(dtype=torch.uint8).detach()
|
||||
|
||||
def decode_latent_to_preview(self, x0, cpudev = torch.device("cpu")):
|
||||
if self.check_use_cached():
|
||||
return self.cached
|
||||
if x0.device == cpudev or not self.use_cuda:
|
||||
return self._decode_latent_to_preview_nocuda(x0)
|
||||
if self.stream is None:
|
||||
self.stream = torch.cuda.Stream()
|
||||
elif not self.stream.query():
|
||||
return self.cached or self.blank
|
||||
work = None
|
||||
if self.prev_work is not None:
|
||||
# We will only arrive here if the stream is ready. Sync just to be safe, should be instant.
|
||||
self.stream.synchronize()
|
||||
work = self.prev_work
|
||||
del self.prev_work
|
||||
# The default stream may be still processing the current step.
|
||||
self.stream.wait_stream(torch.cuda.default_stream())
|
||||
try:
|
||||
torch.cuda.set_stream(self.stream)
|
||||
self.prev_work = self._decode_latent(x0).to(device=cpudev, non_blocking=True)
|
||||
finally:
|
||||
torch.cuda.set_stream(torch.cuda.default_stream())
|
||||
return self.work_to_image(work) if work is not None else self.blank
|
||||
def decode_latent_to_preview(self, x0):
|
||||
use_cached = self.check_use_cached()
|
||||
if x0.device == self.cpudev or not self.use_cuda:
|
||||
return (
|
||||
self.cached if use_cached else self._decode_latent_to_preview_nocuda(x0)
|
||||
)
|
||||
if self.stream is None:
|
||||
self.stream = torch.cuda.Stream()
|
||||
elif not self.stream.query():
|
||||
return self.cached or self.blank
|
||||
work = None
|
||||
if self.prev_work is not None:
|
||||
# We will only arrive here if the stream is ready. Sync just to be safe, should be instant.
|
||||
self.stream.synchronize()
|
||||
work = self.prev_work
|
||||
del self.prev_work
|
||||
result = self.work_to_image(work) if work is not None else self.blank
|
||||
if use_cached:
|
||||
return result
|
||||
# The original stream may be still processing the current step.
|
||||
orig_stream = torch.cuda.current_stream()
|
||||
self.stream.wait_stream(orig_stream)
|
||||
try:
|
||||
torch.cuda.set_stream(self.stream)
|
||||
self.prev_work = self._decode_latent(x0).to(
|
||||
device=self.cpudev,
|
||||
non_blocking=True,
|
||||
)
|
||||
finally:
|
||||
torch.cuda.set_stream(orig_stream)
|
||||
return result
|
||||
|
||||
def work_to_image(self, samples):
|
||||
samples = tuple(np.moveaxis(x, 0, 2) for x in samples.numpy())
|
||||
batch_size = len(samples)
|
||||
height, width, _ = samples[0].shape
|
||||
if batch_size < 2:
|
||||
self.cached = Image.fromarray(samples[0])
|
||||
return self.cached
|
||||
cols = min(math.ceil(batch_size / 2), SETTINGS.btp_max_batch_cols)
|
||||
rows = math.ceil(batch_size / cols)
|
||||
self.cached = result = Image.new("RGB", size = (width * cols, height * rows))
|
||||
for idx in range(batch_size):
|
||||
result.paste(Image.fromarray(samples[idx]),
|
||||
box = ((idx % cols) * width, ((idx // cols) % rows) * height),
|
||||
)
|
||||
self.cached = result
|
||||
return result
|
||||
def calc_cols_rows(self, batch_size, width, height):
|
||||
ratio = width / height
|
||||
cols = min(math.ceil(batch_size / 2), SETTINGS.btp_max_batch_cols)
|
||||
rows = math.ceil(batch_size / cols)
|
||||
return cols, rows
|
||||
|
||||
def work_to_image(self, samples):
|
||||
samples = tuple(np.moveaxis(x, 0, 2) for x in samples.numpy())
|
||||
batch_size = len(samples)
|
||||
height, width, _ = samples[0].shape
|
||||
if batch_size < 2:
|
||||
self.cached = Image.fromarray(samples[0])
|
||||
return self.cached
|
||||
cols, rows = self.calc_cols_rows(batch_size, width, height)
|
||||
|
||||
self.cached = result = Image.new("RGB", size=(width * cols, height * rows))
|
||||
for idx in range(batch_size):
|
||||
result.paste(
|
||||
Image.fromarray(samples[idx]),
|
||||
box=((idx % cols) * width, ((idx // cols) % rows) * height),
|
||||
)
|
||||
self.cached = result
|
||||
return result
|
||||
|
||||
def _decode_latent_to_preview_nocuda(self, x0):
|
||||
return self.work_to_image(self._decode_latent(x0).cpu())
|
||||
|
||||
def _decode_latent_to_preview_nocuda(self, x0):
|
||||
return self.work_to_image(self._decode_latent(x0).cpu())
|
||||
|
||||
latent_preview.TAESDPreviewerImpl = BetterTAESDPreviewer
|
||||
|
||||
+177
-90
@@ -2,118 +2,205 @@
|
||||
# Originally taken from: https://github.com/tfernd/HyperTile/
|
||||
# Modified version of ComfyUI main code
|
||||
# https://github.com/comfyanonymous/ComfyUI/blob/master/comfy_extras/nodes_hypertile.py
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
import random
|
||||
|
||||
import torch
|
||||
|
||||
from einops import rearrange
|
||||
from comfy.ldm.modules.diffusionmodules.openaimodel import forward_timestep_embed, timestep_embedding, th, apply_control
|
||||
|
||||
def get_closest_divisors(hw: int, aspect_ratio: float) -> tuple[int, int]:
|
||||
pairs = [(i, hw // i) for i in range(int(math.sqrt(hw)), 1, -1) if hw % i == 0]
|
||||
pair = min(((i, hw // i) for i in range(2, hw + 1) if hw % i == 0),
|
||||
key=lambda x: abs(x[1] / x[0] - aspect_ratio))
|
||||
pairs.append(pair)
|
||||
res = min(pairs, key=lambda x: max(x) / min(x))
|
||||
return res
|
||||
|
||||
|
||||
def calc_optimal_hw(hw: int, aspect_ratio: float) -> tuple[int, int]:
|
||||
hcand = round(math.sqrt(hw * aspect_ratio))
|
||||
wcand = hw // hcand
|
||||
class HyperTile:
|
||||
def __init__(
|
||||
self,
|
||||
model,
|
||||
seed,
|
||||
tile_size,
|
||||
swap_size,
|
||||
max_depth,
|
||||
scale_depth,
|
||||
start_step,
|
||||
end_step,
|
||||
):
|
||||
self.rng = random.Random()
|
||||
self.rng.seed(seed)
|
||||
self.model = model
|
||||
self.latent_tile_size = max(32, tile_size) // 8
|
||||
self.swap_size = swap_size
|
||||
self.max_depth = max_depth
|
||||
self.scale_depth = scale_depth
|
||||
self.start_step = start_step
|
||||
self.end_step = end_step
|
||||
# Temporary storage for rearranged tensors in the output part
|
||||
self.temp = None
|
||||
|
||||
def patch(self):
|
||||
model = self.model
|
||||
model.set_model_attn1_patch(self.attn1_in)
|
||||
model.set_model_attn1_output_patch(self.attn1_out)
|
||||
return model
|
||||
|
||||
@staticmethod
|
||||
def get_closest_divisors(hw: int, aspect_ratio: float) -> tuple[int, int]:
|
||||
pairs = tuple(
|
||||
(i, hw // i) for i in range(int(math.sqrt(hw)), 1, -1) if hw % i == 0
|
||||
)
|
||||
pair = min(
|
||||
((i, hw // i) for i in range(2, hw + 1) if hw % i == 0),
|
||||
key=lambda x: abs(x[1] / x[0] - aspect_ratio),
|
||||
)
|
||||
pairs.append(pair)
|
||||
return min(pairs, key=lambda x: max(x) / min(x))
|
||||
|
||||
@classmethod
|
||||
def calc_optimal_hw(cls, hw: int, aspect_ratio: float) -> tuple[int, int]:
|
||||
hcand = round(math.sqrt(hw * aspect_ratio))
|
||||
wcand = hw // hcand
|
||||
|
||||
if hcand * wcand == hw:
|
||||
return hcand, wcand
|
||||
|
||||
if hcand * wcand != hw:
|
||||
wcand = round(math.sqrt(hw / aspect_ratio))
|
||||
hcand = hw // wcand
|
||||
|
||||
if hcand * wcand != hw:
|
||||
return get_closest_divisors(hw, aspect_ratio)
|
||||
return cls.get_closest_divisors(hw, aspect_ratio)
|
||||
|
||||
return hcand, wcand
|
||||
return hcand, wcand
|
||||
|
||||
def random_divisor(value: int, min_value: int, /, max_options: int = 1, rand_obj=random.Random()) -> int:
|
||||
min_value = min(min_value, value)
|
||||
def random_divisor(
|
||||
self,
|
||||
value: int,
|
||||
min_value: int,
|
||||
/,
|
||||
max_options: int = 1,
|
||||
) -> int:
|
||||
min_value = min(min_value, value)
|
||||
# All big divisors of value (inclusive)
|
||||
divisors = tuple(i for i in range(min_value, value + 1) if value % i == 0)
|
||||
ns = tuple(value // i for i in divisors[:max_options]) # has at least 1 element
|
||||
idx = self.rng.randint(0, len(ns) - 1) if len(ns) - 1 > 0 else 0
|
||||
return ns[idx]
|
||||
|
||||
# All big divisors of value (inclusive)
|
||||
divisors = [i for i in range(min_value, value + 1) if value % i == 0]
|
||||
def check_timestep(self, extra_options):
|
||||
current_timestep = self.model.model.model_sampling.timestep(
|
||||
extra_options["sigmas"][0],
|
||||
).item()
|
||||
return current_timestep <= self.start_step and current_timestep >= self.end_step
|
||||
|
||||
ns = [value // i for i in divisors[:max_options]] # has at least 1 element
|
||||
def attn1_in(self, q, k, v, extra_options):
|
||||
if not self.check_timestep(extra_options):
|
||||
self.temp = None
|
||||
return q, k, v
|
||||
|
||||
if len(ns) - 1 > 0:
|
||||
idx = rand_obj.randint(0, len(ns) - 1)
|
||||
else:
|
||||
idx = 0
|
||||
model_chans = q.shape[-2]
|
||||
orig_shape = extra_options["original_shape"]
|
||||
|
||||
apply_to = tuple(
|
||||
(orig_shape[-2] / (2**i)) * (orig_shape[-1] / (2**i))
|
||||
for i in range(self.max_depth + 1)
|
||||
)
|
||||
if model_chans not in apply_to:
|
||||
return q, k, v
|
||||
|
||||
aspect_ratio = orig_shape[-1] / orig_shape[-2]
|
||||
|
||||
hw = q.size(1)
|
||||
h, w = (
|
||||
round(math.sqrt(hw * aspect_ratio)),
|
||||
round(math.sqrt(hw / aspect_ratio)),
|
||||
)
|
||||
|
||||
factor = (2 ** apply_to.index(model_chans)) if self.scale_depth else 1
|
||||
|
||||
nh = self.random_divisor(h, self.latent_tile_size * factor, self.swap_size)
|
||||
nw = self.random_divisor(w, self.latent_tile_size * factor, self.swap_size)
|
||||
|
||||
if nh * nw <= 1:
|
||||
return q, k, v
|
||||
|
||||
q = rearrange(
|
||||
q,
|
||||
"b (nh h nw w) c -> (b nh nw) (h w) c",
|
||||
h=h // nh,
|
||||
w=w // nw,
|
||||
nh=nh,
|
||||
nw=nw,
|
||||
)
|
||||
self.temp = (nh, nw, h, w)
|
||||
return q, k, v
|
||||
|
||||
def attn1_out(self, out, _extra_options):
|
||||
if self.temp is None:
|
||||
return out
|
||||
nh, nw, h, w = self.temp
|
||||
self.temp = None
|
||||
out = rearrange(out, "(b nh nw) hw c -> b nh nw hw c", nh=nh, nw=nw)
|
||||
return rearrange(
|
||||
out,
|
||||
"b nh nw (h w) c -> b (nh h nw w) c",
|
||||
h=h // nh,
|
||||
w=w // nw,
|
||||
)
|
||||
|
||||
return ns[idx]
|
||||
|
||||
class HyperTileBleh:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"model": ("MODEL",),
|
||||
"seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
|
||||
"tile_size": ("INT", {"default": 256, "min": 1, "max": 2048}),
|
||||
"swap_size": ("INT", {"default": 2, "min": 1, "max": 128}),
|
||||
"max_depth": ("INT", {"default": 0, "min": 0, "max": 10}),
|
||||
"scale_depth": ("BOOLEAN", {"default": False}),
|
||||
"start_step": ("INT", { "default": 0, "min": 0, "max": 1000, "step": 1, "display": "number" }),
|
||||
"end_step": ("INT", { "default": 1000, "min": 0, "max": 1000, "step": 0.1, }),
|
||||
}}
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"model": ("MODEL",),
|
||||
"seed": ("INT", {"default": 0, "min": 0, "max": 0xFFFFFFFFFFFFFFFF}),
|
||||
"tile_size": ("INT", {"default": 256, "min": 1, "max": 2048}),
|
||||
"swap_size": ("INT", {"default": 2, "min": 1, "max": 128}),
|
||||
"max_depth": ("INT", {"default": 0, "min": 0, "max": 10}),
|
||||
"scale_depth": ("BOOLEAN", {"default": False}),
|
||||
"start_step": (
|
||||
"INT",
|
||||
{
|
||||
"default": 0,
|
||||
"min": 0,
|
||||
"max": 1000,
|
||||
"step": 1,
|
||||
"display": "number",
|
||||
},
|
||||
),
|
||||
"end_step": (
|
||||
"INT",
|
||||
{
|
||||
"default": 1000,
|
||||
"min": 0,
|
||||
"max": 1000,
|
||||
"step": 0.1,
|
||||
},
|
||||
),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("MODEL",)
|
||||
FUNCTION = "patch"
|
||||
|
||||
CATEGORY = "bleh/model_patches"
|
||||
|
||||
def patch(self, model, seed, tile_size, swap_size, max_depth, scale_depth, start_step, end_step):
|
||||
latent_tile_size = max(32, tile_size) // 8
|
||||
temp = None
|
||||
|
||||
rand_obj = random.Random()
|
||||
rand_obj.seed(seed)
|
||||
|
||||
def hypertile_in(q, k, v, extra_options):
|
||||
nonlocal temp
|
||||
|
||||
current_timestep = model.model.model_sampling.timestep(extra_options['sigmas'][0]).item()
|
||||
if current_timestep > start_step or current_timestep < end_step:
|
||||
temp = None
|
||||
return q, k, v
|
||||
|
||||
model_chans = q.shape[-2]
|
||||
orig_shape = extra_options['original_shape']
|
||||
|
||||
apply_to = []
|
||||
for i in range(max_depth + 1):
|
||||
apply_to.append((orig_shape[-2] / (2 ** i)) * (orig_shape[-1] / (2 ** i)))
|
||||
if model_chans not in apply_to:
|
||||
return q, k, v
|
||||
|
||||
aspect_ratio = orig_shape[-1] / orig_shape[-2]
|
||||
|
||||
hw = q.size(1)
|
||||
h, w = round(math.sqrt(hw * aspect_ratio)), round(math.sqrt(hw / aspect_ratio))
|
||||
|
||||
factor = (2 ** apply_to.index(model_chans)) if scale_depth else 1
|
||||
nh = random_divisor(h, latent_tile_size * factor, swap_size, rand_obj)
|
||||
nw = random_divisor(w, latent_tile_size * factor, swap_size, rand_obj)
|
||||
|
||||
if nh * nw > 1:
|
||||
q = rearrange(q, "b (nh h nw w) c -> (b nh nw) (h w) c", h=h // nh, w=w // nw, nh=nh, nw=nw)
|
||||
temp = (nh, nw, h, w)
|
||||
return q, k, v
|
||||
|
||||
def hypertile_out(out, extra_options):
|
||||
nonlocal temp
|
||||
if temp is not None:
|
||||
nh, nw, h, w = temp
|
||||
temp = None
|
||||
out = rearrange(out, "(b nh nw) hw c -> b nh nw hw c", nh=nh, nw=nw)
|
||||
out = rearrange(out, "b nh nw (h w) c -> b (nh h nw w) c", h=h // nh, w=w // nw)
|
||||
return out
|
||||
|
||||
m = model.clone()
|
||||
m.set_model_attn1_patch(hypertile_in)
|
||||
m.set_model_attn1_output_patch(hypertile_out)
|
||||
return (m, )
|
||||
def patch(
|
||||
self,
|
||||
model,
|
||||
seed,
|
||||
tile_size,
|
||||
swap_size,
|
||||
max_depth,
|
||||
scale_depth,
|
||||
start_step,
|
||||
end_step,
|
||||
):
|
||||
return (
|
||||
HyperTile(
|
||||
model.clone(),
|
||||
seed,
|
||||
tile_size,
|
||||
swap_size,
|
||||
max_depth,
|
||||
scale_depth,
|
||||
start_step,
|
||||
end_step,
|
||||
).patch(),
|
||||
)
|
||||
|
||||
+44
-24
@@ -1,29 +1,49 @@
|
||||
import json
|
||||
import os
|
||||
from pathlib import Path
|
||||
|
||||
class Settings:
|
||||
def __init__(self, obj = {}):
|
||||
btp = obj.get("betterTaesdPreviews", None)
|
||||
if btp is None:
|
||||
self.btp_enabled = False
|
||||
else:
|
||||
self.btp_enabled = True
|
||||
self.btp_max_size = btp.get("max_size", 768)
|
||||
self.btp_max_batch = btp.get("max_batch", 4)
|
||||
self.btp_max_batch_cols = btp.get("max_batch_cols", 2)
|
||||
self.btp_throttle_secs = btp.get("throttle_secs", 1)
|
||||
self.btp_use_cuda = btp.get("use_cuda", True)
|
||||
|
||||
SETTINGS = None
|
||||
class Settings:
|
||||
def __init__(self):
|
||||
self.btp_enabled = False
|
||||
|
||||
def update(self, obj):
|
||||
btp = obj.get("betterTaesdPreviews", None)
|
||||
if btp is None:
|
||||
self.btp_enabled = False
|
||||
else:
|
||||
self.btp_enabled = True
|
||||
self.btp_max_size = btp.get("max_size", 768)
|
||||
self.btp_max_batch = btp.get("max_batch", 4)
|
||||
self.btp_max_batch_cols = btp.get("max_batch_cols", 2)
|
||||
self.btp_throttle_secs = btp.get("throttle_secs", 1)
|
||||
self.btp_use_cuda = btp.get("use_cuda", True)
|
||||
|
||||
def get_cfg_path(self, filename):
|
||||
my_path = Path.resolve(Path(__file__).parent)
|
||||
return my_path.parent / filename
|
||||
|
||||
def try_update_from_json(self, filename):
|
||||
import json
|
||||
|
||||
try:
|
||||
with Path.open(self.get_cfg_path(filename)) as fp:
|
||||
self.update(json.load(fp))
|
||||
except OSError:
|
||||
return False
|
||||
|
||||
def try_update_from_yaml(self, filename):
|
||||
try:
|
||||
import yaml
|
||||
|
||||
with Path.open(self.get_cfg_path(filename)) as fp:
|
||||
self.update(yaml.safe_load(fp))
|
||||
except (OSError, ImportError):
|
||||
return False
|
||||
|
||||
|
||||
SETTINGS = Settings()
|
||||
|
||||
|
||||
def load_settings():
|
||||
global SETTINGS
|
||||
my_path = Path(os.path.abspath(os.path.dirname(__file__)))
|
||||
cfg_path = my_path.parent / "blehconfig.json"
|
||||
try:
|
||||
with open(cfg_path, "r") as fp:
|
||||
SETTINGS = Settings(json.load(fp))
|
||||
except OSError:
|
||||
SETTINGS = Settings()
|
||||
return SETTINGS
|
||||
if not SETTINGS.try_update_from_yaml("blehconfig.yaml"):
|
||||
SETTINGS.try_update_from_json("blehconfig.json")
|
||||
return SETTINGS
|
||||
|
||||
@@ -0,0 +1,34 @@
|
||||
[lint]
|
||||
ignore = [
|
||||
"ANN001",
|
||||
"ANN101",
|
||||
"ANN102",
|
||||
"ANN201",
|
||||
"ANN202",
|
||||
"ANN204",
|
||||
"ANN206",
|
||||
"C901",
|
||||
"D100",
|
||||
"D101",
|
||||
"D102",
|
||||
"D103",
|
||||
"D104",
|
||||
"D107",
|
||||
"D211",
|
||||
"D213",
|
||||
"E402",
|
||||
"E501",
|
||||
"EM101",
|
||||
"ERA001",
|
||||
"F403",
|
||||
"F405",
|
||||
"PLR0912",
|
||||
"PLR0913",
|
||||
"PLR0915",
|
||||
"PLR2004",
|
||||
"T201",
|
||||
"TRY003",
|
||||
"N802",
|
||||
"N999",
|
||||
]
|
||||
select = ["ALL"]
|
||||
Reference in New Issue
Block a user