fix error

This commit is contained in:
LouieStark
2024-03-31 19:26:14 +08:00
parent bf53829530
commit e00c23d09a
16 changed files with 15 additions and 23 deletions
+1
View File
@@ -1,4 +1,5 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
"""Concise re-implementation of ``https://github.com/openai/CLIP'' and
``https://github.com/mlfoundations/open_clip''.
"""
@@ -700,7 +700,6 @@ class GeneralConditioner(BaseEmbedder):
if isinstance(emb_out, dict):
for key, val in emb_out.items():
if key in output:
# 重复出现的key必须在(y, crossattn, concat)中,否则raise error
assert key in self.KEY2CATDIM
output[key] = torch.cat([output[key], val],
dim=self.KEY2CATDIM[key])
@@ -715,7 +714,6 @@ class GeneralConditioner(BaseEmbedder):
emb_out = [emb_out]
for emb in emb_out:
# 根据emb的维度,判断该cond归属于 (y, concat, crossattn)中的哪一种
out_key = self.OUTPUT_DIM2KEYS[emb.dim()]
if embedder.ucg_rate > 0.0 and embedder.legacy_ucg_val is None:
+3 -2
View File
@@ -1,4 +1,5 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
from scepter.modules.model.metric.classification import (
AccuracyMetric, EnsembleAccuracyMetric)
from scepter.modules.model.metric.classification import (AccuracyMetric,
EnsembleAccuracyMetric
)
+1 -11
View File
@@ -117,17 +117,7 @@ def crop_back(pred, tar_image, extra_sizes, tar_box_yyxx_crop):
H1, W1, H2, W2, pad1, pad2 = extra_sizes
y1, y2, x1, x2 = tar_box_yyxx_crop
pred = TF.resize(pred, (H2, W2), antialias=True)
# if W1 == H1:
# tar_image[:, y1:y2, x1:x2] = pred
# return tar_image
# if W1 < W2:
# pad1 = int((W2 - W1) / 2)
# pad2 = W2 - W1 - pad1
# pred = pred[:, :,pad1:-pad2]
# else:
# pad1 = int((H2 - H1) / 2)
# pad2 = H2 - H1 - pad1
# pred = pred[:, pad1:-pad2, :]
if W1 < W2:
# pad width
assert H1 == H2 and (pad1 + W1) == (W2 - pad2)
-1
View File
@@ -11,7 +11,6 @@ from scepter.modules.solver.hooks.lr import LrHook
from scepter.modules.solver.hooks.registry import HOOKS
from scepter.modules.solver.hooks.safetensors import SafetensorsHook
from scepter.modules.solver.hooks.sampler import DistSamplerHook
"""
Normally, hooks have priorities, below we recommend priority that runs fine (low score MEANS high priority)
BackwardHook: 0
+1
View File
@@ -1,4 +1,5 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import os.path as osp
from tqdm import tqdm
@@ -339,11 +339,8 @@ class LargenUI(UIBase):
y1, y2, x1, x2 = tar_yyxx_crop
crop_tar_image = tar_image[y1:y2, x1:x2, :]
crop_tar_mask = tar_mask[y1:y2, x1:x2, :]
# 得到pad之前的HW
H1, W1 = crop_tar_image.shape[:2]
# 对tar_mask进行一些预处理
if use_rectangle_mask:
tar_bbox_yyxx = get_bbox_from_mask(crop_tar_mask)
y1, y2, x1, x2 = tar_bbox_yyxx
@@ -351,7 +348,6 @@ class LargenUI(UIBase):
crop_tar_image, pad1, pad2 = pad_to_square(crop_tar_image.astype(np.uint8), pad_value=0)
crop_tar_mask, _, _ = pad_to_square(crop_tar_mask, pad_value=0)
# 得到pad之后的HW
H2, W2 = crop_tar_image.shape[:2]
aug_tar_image = cv2.resize(crop_tar_image.astype(np.uint8), (output_width, output_height))
@@ -379,8 +375,6 @@ class LargenUI(UIBase):
crop_ref_image_i = ref_image[y1:y2, x1:x2, :]
crop_ref_mask_i = ref_mask[y1:y2, x1:x2, :]
# 通过ref_expand_ratio这个值pad subject image, 从而调整clip输入图像中物体的大小
# ref_expand_ratio越大,对应物体越小
h, w = crop_ref_mask_i.shape[:2]
ref_expand_size = int(max(h, w) * 1.02)
pad_op = A.PadIfNeeded(ref_expand_size, ref_expand_size,
@@ -63,7 +63,7 @@ def run_task(cfg):
f'to {hook.prob_interval} according to the setting epoches '
f'interval {ori_interval}')
solver.eval_interval = hook.prob_interval
# size 为无限的时候,使用默认值。
solver.solve()
@@ -1,4 +1,5 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import time
for i in range(180):
@@ -1,4 +1,5 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import json
import os
import subprocess
@@ -1,4 +1,5 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import os
from collections import OrderedDict
@@ -1,4 +1,5 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import gradio as gr
@@ -1,4 +1,5 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import os.path
from scepter.modules.utils.config import Config
@@ -1,4 +1,5 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
def update_2level_dict(d, new_dict):
for first, v in new_dict.items():
if first in d:
@@ -1,4 +1,5 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import re
@@ -1,4 +1,5 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import yaml