Apply no-implicit-optional fixes
This commit is contained in:
parent
39f89026f0
commit
da5477bdbf
@ -14,7 +14,7 @@ class FeatureFusionBlock2d(nn.Module):
|
|||||||
self.deconv = nn.ConvTranspose2d(in_channels=num_features, out_channels=num_features, kernel_size=2, stride=2, padding=0, bias=False) # fmt: off
|
self.deconv = nn.ConvTranspose2d(in_channels=num_features, out_channels=num_features, kernel_size=2, stride=2, padding=0, bias=False) # fmt: off
|
||||||
self.out_conv = nn.Conv2d(in_channels=num_features, out_channels=num_features, kernel_size=1, stride=1, padding=0, bias=True) # fmt: off
|
self.out_conv = nn.Conv2d(in_channels=num_features, out_channels=num_features, kernel_size=1, stride=1, padding=0, bias=True) # fmt: off
|
||||||
|
|
||||||
def __call__(self, x0: mx.array, x1: mx.array = None) -> mx.array:
|
def __call__(self, x0: mx.array, x1: mx.array | None = None) -> mx.array:
|
||||||
x = x0
|
x = x0
|
||||||
if x1 is not None:
|
if x1 is not None:
|
||||||
res = self.resnet1(x1)
|
res = self.resnet1(x1)
|
||||||
|
|||||||
@ -21,7 +21,7 @@ class UpSampleBlock(nn.Module):
|
|||||||
) # fmt: off
|
) # fmt: off
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _create_layers(dim_in: int, dim_out: int, upsample_layers: int, dim_int: int = None) -> list[nn.Module]:
|
def _create_layers(dim_in: int, dim_out: int, upsample_layers: int, dim_int: int | None = None) -> list[nn.Module]:
|
||||||
if dim_int is None:
|
if dim_int is None:
|
||||||
dim_int = dim_out
|
dim_int = dim_out
|
||||||
|
|
||||||
|
|||||||
@ -1,5 +1,4 @@
|
|||||||
import importlib
|
import importlib
|
||||||
import typing as t
|
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
|
||||||
import mlx.core as mx
|
import mlx.core as mx
|
||||||
@ -78,7 +77,7 @@ class GeneratedImage:
|
|||||||
|
|
||||||
def save(
|
def save(
|
||||||
self,
|
self,
|
||||||
path: t.Union[str, Path],
|
path: str | Path,
|
||||||
export_json_metadata: bool = False,
|
export_json_metadata: bool = False,
|
||||||
overwrite: bool = False,
|
overwrite: bool = False,
|
||||||
) -> None:
|
) -> None:
|
||||||
|
|||||||
@ -1,6 +1,5 @@
|
|||||||
import json
|
import json
|
||||||
import logging
|
import logging
|
||||||
import typing as t
|
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
|
||||||
import mlx.core as mx
|
import mlx.core as mx
|
||||||
@ -59,7 +58,7 @@ class ImageUtil:
|
|||||||
)
|
)
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def to_composite_image(generated_images: t.List[GeneratedImage]) -> PIL.Image.Image:
|
def to_composite_image(generated_images: list[GeneratedImage]) -> PIL.Image.Image:
|
||||||
# stitch horizontally
|
# stitch horizontally
|
||||||
total_width = sum(gen_img.image.width for gen_img in generated_images)
|
total_width = sum(gen_img.image.width for gen_img in generated_images)
|
||||||
max_height = max(gen_img.image.height for gen_img in generated_images)
|
max_height = max(gen_img.image.height for gen_img in generated_images)
|
||||||
@ -119,7 +118,7 @@ class ImageUtil:
|
|||||||
@staticmethod
|
@staticmethod
|
||||||
def expand_image(
|
def expand_image(
|
||||||
image: PIL.Image.Image,
|
image: PIL.Image.Image,
|
||||||
box_values: AbsoluteBoxValues = None,
|
box_values: AbsoluteBoxValues | None = None,
|
||||||
top: int | str = 0,
|
top: int | str = 0,
|
||||||
right: int | str = 0,
|
right: int | str = 0,
|
||||||
bottom: int | str = 0,
|
bottom: int | str = 0,
|
||||||
@ -161,7 +160,7 @@ class ImageUtil:
|
|||||||
orig_height: int,
|
orig_height: int,
|
||||||
border_color: tuple,
|
border_color: tuple,
|
||||||
content_color: tuple,
|
content_color: tuple,
|
||||||
box_values: AbsoluteBoxValues = None,
|
box_values: AbsoluteBoxValues | None = None,
|
||||||
top: int | str = 0,
|
top: int | str = 0,
|
||||||
right: int | str = 0,
|
right: int | str = 0,
|
||||||
bottom: int | str = 0,
|
bottom: int | str = 0,
|
||||||
@ -204,7 +203,7 @@ class ImageUtil:
|
|||||||
@staticmethod
|
@staticmethod
|
||||||
def save_image(
|
def save_image(
|
||||||
image: PIL.Image.Image,
|
image: PIL.Image.Image,
|
||||||
path: t.Union[str, Path],
|
path: str | Path,
|
||||||
metadata: dict | None = None,
|
metadata: dict | None = None,
|
||||||
export_json_metadata: bool = False,
|
export_json_metadata: bool = False,
|
||||||
overwrite: bool = False,
|
overwrite: bool = False,
|
||||||
|
|||||||
@ -7,9 +7,9 @@ from huggingface_hub import snapshot_download
|
|||||||
class WeightHandlerLoRAHuggingFace:
|
class WeightHandlerLoRAHuggingFace:
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def download_loras(
|
def download_loras(
|
||||||
lora_names: list[str] = None,
|
lora_names: list[str] | None = None,
|
||||||
repo_id: str = None,
|
repo_id: str | None = None,
|
||||||
cache_dir: str = None,
|
cache_dir: str | None = None,
|
||||||
) -> list[str]:
|
) -> list[str]:
|
||||||
if repo_id is None:
|
if repo_id is None:
|
||||||
return []
|
return []
|
||||||
@ -30,7 +30,7 @@ class WeightHandlerLoRAHuggingFace:
|
|||||||
def _download_lora(
|
def _download_lora(
|
||||||
repo_id: str,
|
repo_id: str,
|
||||||
lora_name: str,
|
lora_name: str,
|
||||||
cache_dir: str = None,
|
cache_dir: str | None = None,
|
||||||
) -> str:
|
) -> str:
|
||||||
# Create cache directory if it doesn't exist
|
# Create cache directory if it doesn't exist
|
||||||
if cache_dir is None:
|
if cache_dir is None:
|
||||||
|
|||||||
@ -17,8 +17,8 @@ class ImageGeneratorDepthTestHelper:
|
|||||||
prompt: str,
|
prompt: str,
|
||||||
steps: int,
|
steps: int,
|
||||||
seed: int,
|
seed: int,
|
||||||
height: int = None,
|
height: int | None = None,
|
||||||
width: int = None,
|
width: int | None = None,
|
||||||
image_path: str | None = None,
|
image_path: str | None = None,
|
||||||
depth_image_path: str | None = None,
|
depth_image_path: str | None = None,
|
||||||
):
|
):
|
||||||
|
|||||||
@ -18,8 +18,8 @@ class ImageGeneratorInContextTestHelper:
|
|||||||
prompt: str,
|
prompt: str,
|
||||||
steps: int,
|
steps: int,
|
||||||
seed: int,
|
seed: int,
|
||||||
height: int = None,
|
height: int | None = None,
|
||||||
width: int = None,
|
width: int | None = None,
|
||||||
image_path: str | None = None,
|
image_path: str | None = None,
|
||||||
lora_style: str | None = None,
|
lora_style: str | None = None,
|
||||||
lora_paths: list[str] | None = None,
|
lora_paths: list[str] | None = None,
|
||||||
|
|||||||
@ -16,8 +16,8 @@ class ImageGeneratorTestHelper:
|
|||||||
prompt: str,
|
prompt: str,
|
||||||
steps: int,
|
steps: int,
|
||||||
seed: int,
|
seed: int,
|
||||||
height: int = None,
|
height: int | None = None,
|
||||||
width: int = None,
|
width: int | None = None,
|
||||||
image_path: str | None = None,
|
image_path: str | None = None,
|
||||||
image_strength: float | None = None,
|
image_strength: float | None = None,
|
||||||
lora_paths: list[str] | None = None,
|
lora_paths: list[str] | None = None,
|
||||||
|
|||||||
Loading…
Reference in New Issue
Block a user