fix bug of Lama
This commit is contained in:
@@ -3,9 +3,9 @@ import cv2
|
||||
import torch
|
||||
from loguru import logger
|
||||
|
||||
from iopaint.const import INSTRUCT_PIX2PIX_NAME
|
||||
from ..const import INSTRUCT_PIX2PIX_NAME
|
||||
from .base import DiffusionInpaintModel
|
||||
from iopaint.schema import InpaintRequest
|
||||
from ..schema import InpaintRequest
|
||||
from .utils import get_torch_dtype, enable_low_mem, is_local_files_only
|
||||
|
||||
|
||||
|
||||
@@ -3,9 +3,9 @@ import cv2
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
from iopaint.const import KANDINSKY22_NAME
|
||||
from ..const import KANDINSKY22_NAME
|
||||
from .base import DiffusionInpaintModel
|
||||
from iopaint.schema import InpaintRequest
|
||||
from ..schema import InpaintRequest
|
||||
from .utils import get_torch_dtype, enable_low_mem, is_local_files_only
|
||||
|
||||
|
||||
|
||||
@@ -4,13 +4,13 @@ import cv2
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
from iopaint.helper import (
|
||||
from ..helper import (
|
||||
norm_img,
|
||||
get_cache_path_by_url,
|
||||
load_jit_model,
|
||||
download_model,
|
||||
)
|
||||
from iopaint.schema import InpaintRequest
|
||||
from ..schema import InpaintRequest
|
||||
from .base import InpaintModel
|
||||
|
||||
LAMA_MODEL_URL = os.environ.get(
|
||||
|
||||
@@ -7,11 +7,11 @@ from loguru import logger
|
||||
from .base import InpaintModel
|
||||
from .ddim_sampler import DDIMSampler
|
||||
from .plms_sampler import PLMSSampler
|
||||
from iopaint.schema import InpaintRequest, LDMSampler
|
||||
from ..schema import InpaintRequest, LDMSampler
|
||||
|
||||
torch.manual_seed(42)
|
||||
import torch.nn as nn
|
||||
from iopaint.helper import (
|
||||
from ..helper import (
|
||||
download_model,
|
||||
norm_img,
|
||||
get_cache_path_by_url,
|
||||
|
||||
@@ -7,9 +7,9 @@ import torch
|
||||
import time
|
||||
from loguru import logger
|
||||
|
||||
from iopaint.helper import get_cache_path_by_url, load_jit_model, download_model
|
||||
from ..helper import get_cache_path_by_url, load_jit_model, download_model
|
||||
from .base import InpaintModel
|
||||
from iopaint.schema import InpaintRequest
|
||||
from ..schema import InpaintRequest
|
||||
|
||||
|
||||
MANGA_INPAINTOR_MODEL_URL = os.environ.get(
|
||||
|
||||
@@ -8,13 +8,13 @@ import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
import torch.utils.checkpoint as checkpoint
|
||||
|
||||
from iopaint.helper import (
|
||||
from ..helper import (
|
||||
load_model,
|
||||
get_cache_path_by_url,
|
||||
norm_img,
|
||||
download_model,
|
||||
)
|
||||
from iopaint.schema import InpaintRequest
|
||||
from ..schema import InpaintRequest
|
||||
from .base import InpaintModel
|
||||
from .utils import (
|
||||
setup_filter,
|
||||
|
||||
@@ -3,7 +3,7 @@ import os
|
||||
import cv2
|
||||
import torch
|
||||
|
||||
from iopaint.helper import (
|
||||
from ..helper import (
|
||||
load_jit_model,
|
||||
download_model,
|
||||
get_cache_path_by_url,
|
||||
@@ -12,7 +12,7 @@ from iopaint.helper import (
|
||||
norm_img,
|
||||
)
|
||||
from .base import InpaintModel
|
||||
from iopaint.schema import InpaintRequest
|
||||
from ..schema import InpaintRequest
|
||||
|
||||
MIGAN_MODEL_URL = os.environ.get(
|
||||
"MIGAN_MODEL_URL",
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
import cv2
|
||||
from .base import InpaintModel
|
||||
from iopaint.schema import InpaintRequest
|
||||
from ..schema import InpaintRequest
|
||||
|
||||
flag_map = {"INPAINT_NS": cv2.INPAINT_NS, "INPAINT_TELEA": cv2.INPAINT_TELEA}
|
||||
|
||||
|
||||
@@ -4,9 +4,9 @@ import cv2
|
||||
import torch
|
||||
from loguru import logger
|
||||
|
||||
from iopaint.helper import decode_base64_to_image
|
||||
from ..helper import decode_base64_to_image
|
||||
from .base import DiffusionInpaintModel
|
||||
from iopaint.schema import InpaintRequest
|
||||
from ..schema import InpaintRequest
|
||||
from .utils import get_torch_dtype, enable_low_mem, is_local_files_only
|
||||
|
||||
|
||||
|
||||
@@ -12,7 +12,7 @@ from ..utils import (
|
||||
enable_low_mem,
|
||||
is_local_files_only,
|
||||
)
|
||||
from iopaint.schema import InpaintRequest
|
||||
from ...schema import InpaintRequest
|
||||
from .powerpaint_tokenizer import add_task_to_prompt
|
||||
from ...const import POWERPAINT_NAME
|
||||
|
||||
|
||||
@@ -5,7 +5,7 @@ import random
|
||||
from typing import Any, List, Optional, Union
|
||||
from transformers import CLIPTokenizer
|
||||
|
||||
from iopaint.schema import PowerPaintTask
|
||||
from ...schema import PowerPaintTask
|
||||
|
||||
|
||||
def add_task_to_prompt(prompt, negative_prompt, task: PowerPaintTask):
|
||||
|
||||
@@ -12,7 +12,7 @@ from .utils import (
|
||||
enable_low_mem,
|
||||
is_local_files_only,
|
||||
)
|
||||
from iopaint.schema import InpaintRequest, ModelType
|
||||
from ..schema import InpaintRequest, ModelType
|
||||
|
||||
|
||||
class SD(DiffusionInpaintModel):
|
||||
|
||||
@@ -6,7 +6,7 @@ import torch
|
||||
from diffusers import AutoencoderKL
|
||||
from loguru import logger
|
||||
|
||||
from iopaint.schema import InpaintRequest, ModelType
|
||||
from ..schema import InpaintRequest, ModelType
|
||||
|
||||
from .base import DiffusionInpaintModel
|
||||
from .helper.cpu_text_encoder import CPUTextEncoderWrapper
|
||||
|
||||
@@ -25,7 +25,7 @@ from diffusers import (
|
||||
)
|
||||
from loguru import logger
|
||||
|
||||
from iopaint.schema import SDSampler
|
||||
from ..schema import SDSampler
|
||||
from torch import conv2d, conv_transpose2d
|
||||
|
||||
|
||||
|
||||
@@ -5,8 +5,8 @@ import cv2
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
|
||||
from iopaint.helper import get_cache_path_by_url, load_jit_model, download_model
|
||||
from iopaint.schema import InpaintRequest
|
||||
from ..helper import get_cache_path_by_url, load_jit_model, download_model
|
||||
from ..schema import InpaintRequest
|
||||
import numpy as np
|
||||
|
||||
from .base import InpaintModel
|
||||
|
||||
@@ -5,9 +5,9 @@ import torch.nn.functional as F
|
||||
import numpy as np
|
||||
from PIL import Image
|
||||
|
||||
from iopaint.helper import load_model
|
||||
from iopaint.plugins.base_plugin import BasePlugin
|
||||
from iopaint.schema import RunPluginRequest
|
||||
from ..helper import load_model
|
||||
from ..plugins.base_plugin import BasePlugin
|
||||
from ..schema import RunPluginRequest
|
||||
|
||||
|
||||
class REBNCONV(nn.Module):
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
from loguru import logger
|
||||
import numpy as np
|
||||
|
||||
from iopaint.schema import RunPluginRequest
|
||||
from ..schema import RunPluginRequest
|
||||
|
||||
|
||||
class BasePlugin:
|
||||
|
||||
@@ -2,9 +2,9 @@ import cv2
|
||||
import numpy as np
|
||||
from loguru import logger
|
||||
|
||||
from iopaint.helper import download_model
|
||||
from iopaint.plugins.base_plugin import BasePlugin
|
||||
from iopaint.schema import RunPluginRequest
|
||||
from ..helper import download_model
|
||||
from ..plugins.base_plugin import BasePlugin
|
||||
from ..schema import RunPluginRequest
|
||||
|
||||
|
||||
class GFPGANPlugin(BasePlugin):
|
||||
|
||||
@@ -5,11 +5,11 @@ import numpy as np
|
||||
import torch
|
||||
from loguru import logger
|
||||
|
||||
from iopaint.helper import download_model
|
||||
from iopaint.plugins.base_plugin import BasePlugin
|
||||
from iopaint.plugins.segment_anything import SamPredictor, sam_model_registry
|
||||
from iopaint.plugins.segment_anything.predictor_hq import SamHQPredictor
|
||||
from iopaint.schema import RunPluginRequest
|
||||
from ..helper import download_model
|
||||
from ..plugins.base_plugin import BasePlugin
|
||||
from ..plugins.segment_anything import SamPredictor, sam_model_registry
|
||||
from ..plugins.segment_anything.predictor_hq import SamHQPredictor
|
||||
from ..schema import RunPluginRequest
|
||||
|
||||
# 从小到大
|
||||
SEGMENT_ANYTHING_MODELS = {
|
||||
|
||||
@@ -3,9 +3,9 @@ import numpy as np
|
||||
import torch
|
||||
from loguru import logger
|
||||
|
||||
from iopaint.helper import download_model
|
||||
from iopaint.plugins.base_plugin import BasePlugin
|
||||
from iopaint.schema import RunPluginRequest, RealESRGANModel
|
||||
from ..helper import download_model
|
||||
from ..plugins.base_plugin import BasePlugin
|
||||
from ..schema import RunPluginRequest, RealESRGANModel
|
||||
|
||||
|
||||
class RealESRGANUpscaler(BasePlugin):
|
||||
|
||||
@@ -4,8 +4,8 @@ import numpy as np
|
||||
from loguru import logger
|
||||
from torch.hub import get_dir
|
||||
|
||||
from iopaint.plugins.base_plugin import BasePlugin
|
||||
from iopaint.schema import RunPluginRequest, RemoveBGModel
|
||||
from ..plugins.base_plugin import BasePlugin
|
||||
from ..schema import RunPluginRequest, RemoveBGModel
|
||||
|
||||
|
||||
class RemoveBG(BasePlugin):
|
||||
|
||||
@@ -2,9 +2,9 @@ import cv2
|
||||
import numpy as np
|
||||
from loguru import logger
|
||||
|
||||
from iopaint.helper import download_model
|
||||
from iopaint.plugins.base_plugin import BasePlugin
|
||||
from iopaint.schema import RunPluginRequest
|
||||
from ..helper import download_model
|
||||
from ..plugins.base_plugin import BasePlugin
|
||||
from ..schema import RunPluginRequest
|
||||
|
||||
|
||||
class RestoreFormerPlugin(BasePlugin):
|
||||
|
||||
@@ -8,7 +8,7 @@ import torch
|
||||
|
||||
from functools import partial
|
||||
|
||||
from iopaint.plugins.segment_anything.modeling.tiny_vit_sam import TinyViT
|
||||
from ...plugins.segment_anything.modeling.tiny_vit_sam import TinyViT
|
||||
|
||||
from .modeling import (
|
||||
ImageEncoderViT,
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
import cv2
|
||||
from iopaint.helper import adjust_mask
|
||||
from iopaint.tests.utils import current_dir, save_dir
|
||||
from ..helper import adjust_mask
|
||||
from ..tests.utils import current_dir, save_dir
|
||||
|
||||
mask_p = current_dir / "overture-creations-5sI6fQgYIuo_mask.png"
|
||||
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
import os
|
||||
|
||||
from iopaint.tests.utils import check_device, get_config, assert_equal
|
||||
from ..tests.utils import check_device, get_config, assert_equal
|
||||
|
||||
os.environ["PYTORCH_ENABLE_MPS_FALLBACK"] = "1"
|
||||
from pathlib import Path
|
||||
@@ -8,8 +8,8 @@ from pathlib import Path
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from iopaint.model_manager import ModelManager
|
||||
from iopaint.schema import HDStrategy
|
||||
from ..model_manager import ModelManager
|
||||
from ..schema import HDStrategy
|
||||
|
||||
current_dir = Path(__file__).parent.absolute().resolve()
|
||||
save_dir = current_dir / "result"
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
import os
|
||||
|
||||
from iopaint.const import SD_CONTROLNET_CHOICES
|
||||
from iopaint.tests.utils import current_dir, check_device, get_config, assert_equal
|
||||
from ..const import SD_CONTROLNET_CHOICES
|
||||
from ..tests.utils import current_dir, check_device, get_config, assert_equal
|
||||
|
||||
os.environ["PYTORCH_ENABLE_MPS_FALLBACK"] = "1"
|
||||
from pathlib import Path
|
||||
@@ -9,8 +9,8 @@ from pathlib import Path
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from iopaint.model_manager import ModelManager
|
||||
from iopaint.schema import HDStrategy, SDSampler
|
||||
from ..model_manager import ModelManager
|
||||
from ..schema import HDStrategy, SDSampler
|
||||
|
||||
|
||||
model_name = "runwayml/stable-diffusion-inpainting"
|
||||
|
||||
@@ -3,9 +3,9 @@ from pathlib import Path
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from iopaint.model_manager import ModelManager
|
||||
from iopaint.schema import HDStrategy
|
||||
from iopaint.tests.utils import get_config, check_device, assert_equal, current_dir
|
||||
from ..model_manager import ModelManager
|
||||
from ..schema import HDStrategy
|
||||
from ..tests.utils import get_config, check_device, assert_equal, current_dir
|
||||
|
||||
model_name = "timbrooks/instruct-pix2pix"
|
||||
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
from iopaint.helper import load_img
|
||||
from iopaint.tests.utils import current_dir
|
||||
from ..helper import load_img
|
||||
from ..tests.utils import current_dir
|
||||
|
||||
png_img_p = current_dir / "image.png"
|
||||
jpg_img_p = current_dir / "bunny.jpeg"
|
||||
|
||||
@@ -2,15 +2,15 @@ import os
|
||||
|
||||
from loguru import logger
|
||||
|
||||
from iopaint.tests.utils import check_device, get_config, assert_equal, current_dir
|
||||
from ..tests.utils import check_device, get_config, assert_equal, current_dir
|
||||
|
||||
os.environ["PYTORCH_ENABLE_MPS_FALLBACK"] = "1"
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from iopaint.model_manager import ModelManager
|
||||
from iopaint.schema import HDStrategy, SDSampler, FREEUConfig
|
||||
from ..model_manager import ModelManager
|
||||
from ..schema import HDStrategy, SDSampler, FREEUConfig
|
||||
|
||||
|
||||
@pytest.mark.parametrize("device", ["cuda", "mps"])
|
||||
|
||||
@@ -1,9 +1,9 @@
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from iopaint.model_manager import ModelManager
|
||||
from iopaint.schema import SDSampler, HDStrategy
|
||||
from iopaint.tests.utils import check_device, get_config, assert_equal, current_dir
|
||||
from ..model_manager import ModelManager
|
||||
from ..schema import SDSampler, HDStrategy
|
||||
from ..tests.utils import check_device, get_config, assert_equal, current_dir
|
||||
|
||||
|
||||
@pytest.mark.parametrize("device", ["cuda", "mps"])
|
||||
|
||||
@@ -1,9 +1,9 @@
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from iopaint.model_manager import ModelManager
|
||||
from iopaint.schema import HDStrategy, LDMSampler
|
||||
from iopaint.tests.utils import assert_equal, get_config, current_dir, check_device
|
||||
from ..model_manager import ModelManager
|
||||
from ..schema import HDStrategy, LDMSampler
|
||||
from ..tests.utils import assert_equal, get_config, current_dir, check_device
|
||||
|
||||
|
||||
@pytest.mark.parametrize("device", ["cuda", "mps", "cpu"])
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
def test_load_model():
|
||||
from iopaint.plugins import InteractiveSeg
|
||||
from iopaint.model_manager import ModelManager
|
||||
from ..plugins import InteractiveSeg
|
||||
from ..model_manager import ModelManager
|
||||
|
||||
interactive_seg_model = InteractiveSeg("vit_l", "cpu")
|
||||
|
||||
|
||||
@@ -1,12 +1,12 @@
|
||||
import os
|
||||
|
||||
from iopaint.schema import InpaintRequest
|
||||
from ..schema import InpaintRequest
|
||||
|
||||
os.environ["PYTORCH_ENABLE_MPS_FALLBACK"] = "1"
|
||||
|
||||
import torch
|
||||
|
||||
from iopaint.model_manager import ModelManager
|
||||
from ..model_manager import ModelManager
|
||||
|
||||
|
||||
def test_model_switch():
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
import os
|
||||
|
||||
from iopaint.tests.utils import current_dir, check_device
|
||||
from ..tests.utils import current_dir, check_device
|
||||
|
||||
os.environ["PYTORCH_ENABLE_MPS_FALLBACK"] = "1"
|
||||
from pathlib import Path
|
||||
@@ -8,9 +8,9 @@ from pathlib import Path
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from iopaint.model_manager import ModelManager
|
||||
from iopaint.schema import HDStrategy, SDSampler
|
||||
from iopaint.tests.test_model import get_config, assert_equal
|
||||
from ..model_manager import ModelManager
|
||||
from ..schema import HDStrategy, SDSampler
|
||||
from ..tests.test_model import get_config, assert_equal
|
||||
|
||||
|
||||
@pytest.mark.parametrize("name", ["runwayml/stable-diffusion-inpainting"])
|
||||
|
||||
@@ -2,9 +2,9 @@ import cv2
|
||||
import pytest
|
||||
from PIL import Image
|
||||
|
||||
from iopaint.model_manager import ModelManager
|
||||
from iopaint.schema import HDStrategy
|
||||
from iopaint.tests.utils import (
|
||||
from ..model_manager import ModelManager
|
||||
from ..schema import HDStrategy
|
||||
from ..tests.utils import (
|
||||
current_dir,
|
||||
get_config,
|
||||
get_data,
|
||||
|
||||
@@ -3,17 +3,17 @@ import os
|
||||
import time
|
||||
from PIL import Image
|
||||
|
||||
from iopaint.helper import encode_pil_to_base64, gen_frontend_mask
|
||||
from iopaint.plugins.anime_seg import AnimeSeg
|
||||
from iopaint.schema import RunPluginRequest, RemoveBGModel, InteractiveSegModel
|
||||
from iopaint.tests.utils import check_device, current_dir, save_dir
|
||||
from ..helper import encode_pil_to_base64, gen_frontend_mask
|
||||
from ..plugins.anime_seg import AnimeSeg
|
||||
from ..schema import RunPluginRequest, RemoveBGModel, InteractiveSegModel
|
||||
from ..tests.utils import check_device, current_dir, save_dir
|
||||
|
||||
os.environ["PYTORCH_ENABLE_MPS_FALLBACK"] = "1"
|
||||
|
||||
import cv2
|
||||
import pytest
|
||||
|
||||
from iopaint.plugins import (
|
||||
from ..plugins import (
|
||||
RemoveBG,
|
||||
RealESRGANUpscaler,
|
||||
GFPGANPlugin,
|
||||
|
||||
@@ -5,7 +5,7 @@ from typing import List
|
||||
|
||||
from PIL import Image
|
||||
|
||||
from iopaint.helper import pil_to_bytes, load_img
|
||||
from ..helper import pil_to_bytes, load_img
|
||||
|
||||
current_dir = Path(__file__).parent.absolute().resolve()
|
||||
|
||||
|
||||
@@ -2,7 +2,7 @@ import os
|
||||
|
||||
from loguru import logger
|
||||
|
||||
from iopaint.tests.utils import check_device, get_config, assert_equal
|
||||
from ..tests.utils import check_device, get_config, assert_equal
|
||||
|
||||
os.environ["PYTORCH_ENABLE_MPS_FALLBACK"] = "1"
|
||||
from pathlib import Path
|
||||
@@ -10,8 +10,8 @@ from pathlib import Path
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from iopaint.model_manager import ModelManager
|
||||
from iopaint.schema import HDStrategy, SDSampler, FREEUConfig
|
||||
from ..model_manager import ModelManager
|
||||
from ..schema import HDStrategy, SDSampler, FREEUConfig
|
||||
|
||||
current_dir = Path(__file__).parent.absolute().resolve()
|
||||
save_dir = current_dir / "result"
|
||||
|
||||
@@ -1,15 +1,15 @@
|
||||
import os
|
||||
|
||||
from iopaint.tests.utils import check_device, current_dir
|
||||
from ..tests.utils import check_device, current_dir
|
||||
|
||||
os.environ["PYTORCH_ENABLE_MPS_FALLBACK"] = "1"
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from iopaint.model_manager import ModelManager
|
||||
from iopaint.schema import HDStrategy, SDSampler, FREEUConfig
|
||||
from iopaint.tests.test_model import get_config, assert_equal
|
||||
from ..model_manager import ModelManager
|
||||
from ..schema import HDStrategy, SDSampler, FREEUConfig
|
||||
from ..tests.test_model import get_config, assert_equal
|
||||
|
||||
|
||||
@pytest.mark.parametrize("device", ["cuda", "mps"])
|
||||
|
||||
@@ -3,8 +3,8 @@ import cv2
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from iopaint.helper import encode_pil_to_base64
|
||||
from iopaint.schema import LDMSampler, HDStrategy, InpaintRequest, SDSampler
|
||||
from ..helper import encode_pil_to_base64
|
||||
from ..schema import LDMSampler, HDStrategy, InpaintRequest, SDSampler
|
||||
from PIL import Image
|
||||
|
||||
current_dir = Path(__file__).parent.absolute().resolve()
|
||||
|
||||
Reference in New Issue
Block a user