fix error
This commit is contained in:
@@ -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:
|
||||
|
||||
@@ -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
|
||||
)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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,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
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user