v1.0.3 update
This commit is contained in:
@@ -20,7 +20,7 @@ SOLVER:
|
||||
#
|
||||
FILE_SYSTEM:
|
||||
NAME: "ModelscopeFs"
|
||||
TEMP_DIR: "./cache/data"
|
||||
TEMP_DIR: "./cache/cache_data"
|
||||
#
|
||||
TUNER:
|
||||
-
|
||||
|
||||
@@ -11,7 +11,7 @@ SOLVER:
|
||||
# NUM_FOLDS DESCRIPTION: Num folds for training. TYPE: int default: 0
|
||||
NUM_FOLDS: 1
|
||||
# WORK_DIR DESCRIPTION: Save dir of the training log or model. TYPE: str default: ''
|
||||
WORK_DIR: ./exp12/
|
||||
WORK_DIR: ./cache/save_data/example/
|
||||
LOG_FILE: std_log.txt
|
||||
# EVAL_INTERVAL DESCRIPTION: Eval the model interval. TYPE: int default: 1
|
||||
EVAL_INTERVAL: 1
|
||||
@@ -102,7 +102,7 @@ SOLVER:
|
||||
# DATASET DESCRIPTION: the public dataset name TYPE: str default: 'cifar10'
|
||||
DATASET: cifar10
|
||||
# DATA_ROOT DESCRIPTION: the download data save path TYPE: str default: ''
|
||||
DATA_ROOT: ./local_data/cifar10
|
||||
DATA_ROOT: ./cache/cache_data/cifar10
|
||||
# MODE DESCRIPTION: test TYPE: str default: test
|
||||
MODE: test
|
||||
# PIN_MEMORY DESCRIPTION: pin_memory for data loader TYPE: bool default: False
|
||||
|
||||
@@ -0,0 +1,223 @@
|
||||
ENV:
|
||||
BACKEND: nccl
|
||||
#
|
||||
SOLVER:
|
||||
NAME: LatentDiffusionSolver
|
||||
RESUME_FROM:
|
||||
LOAD_MODEL_ONLY: True
|
||||
USE_FSDP: False
|
||||
SHARDING_STRATEGY:
|
||||
USE_AMP: True
|
||||
DTYPE: float16
|
||||
CHANNELS_LAST: True
|
||||
MAX_STEPS: 1000
|
||||
MAX_EPOCHS: -1
|
||||
NUM_FOLDS: 1
|
||||
ACCU_STEP: 1
|
||||
EVAL_INTERVAL: 100
|
||||
RESCALE_LR: False
|
||||
#
|
||||
WORK_DIR: ./cache/save_data/dit_pixart_alpha_1024_lora
|
||||
LOG_FILE: std_log.txt
|
||||
#
|
||||
FILE_SYSTEM:
|
||||
NAME: "ModelscopeFs"
|
||||
TEMP_DIR: "./cache/cache_data"
|
||||
#
|
||||
FREEZE:
|
||||
#
|
||||
TUNER:
|
||||
- NAME: SwiftLoRA
|
||||
R: 128
|
||||
LORA_ALPHA: 128
|
||||
LORA_DROPOUT: 0.0
|
||||
BIAS: "none"
|
||||
TARGET_MODULES: "model.*(.q|.k|.v|.o|mlp.fc1|mlp.fc2)$"
|
||||
#
|
||||
MODEL:
|
||||
NAME: LatentDiffusionPixart
|
||||
PARAMETERIZATION: eps
|
||||
TIMESTEPS: 1000
|
||||
MIN_SNR_GAMMA:
|
||||
ZERO_TERMINAL_SNR: False
|
||||
PRETRAINED_MODEL:
|
||||
IGNORE_KEYS: [ ]
|
||||
SCALE_FACTOR: 0.18215
|
||||
SIZE_FACTOR: 8
|
||||
DECODER_BIAS: 0.5
|
||||
DEFAULT_N_PROMPT:
|
||||
SCHEDULE_ARGS:
|
||||
"NAME": "linear"
|
||||
"BETA_MIN": 0.0001
|
||||
"BETA_MAX": 0.02
|
||||
USE_EMA: False
|
||||
LOAD_REFINER: False
|
||||
#
|
||||
DIFFUSION_MODEL:
|
||||
NAME: PixArt
|
||||
PRETRAINED_MODEL: ms://AI-ModelScope/PixArt-alpha@PixArt-XL-2-1024-MS.pth
|
||||
INPUT_SIZE: 128
|
||||
PATCH_SIZE: 2
|
||||
IN_CHANNELS: 4
|
||||
HIDDEN_SIZE: 1152
|
||||
DEPTH: 28
|
||||
NUM_HEADS: 16
|
||||
MLP_RATIO: 4.0
|
||||
CLASS_DROPOUT_PROB: 0.1
|
||||
PRED_SIGMA: True
|
||||
DROP_PATH: 0.0
|
||||
WINDOW_DIZE: 0
|
||||
USE_REL_POS: False
|
||||
CAPTION_CHANNELS: 4096
|
||||
LEWEI_SCALE: 2
|
||||
MODEL_MAX_LENGTH: 120
|
||||
#
|
||||
FIRST_STAGE_MODEL:
|
||||
NAME: AutoencoderKL
|
||||
PRETRAINED_MODEL: ms://AI-ModelScope/stable-diffusion-2-base@512-base-ema.safetensors
|
||||
EMBED_DIM: 4
|
||||
IGNORE_KEYS: [ ]
|
||||
BATCH_SIZE: 1
|
||||
#
|
||||
ENCODER:
|
||||
NAME: Encoder
|
||||
CH: 128
|
||||
OUT_CH: 3
|
||||
NUM_RES_BLOCKS: 2
|
||||
IN_CHANNELS: 3
|
||||
ATTN_RESOLUTIONS: [ ]
|
||||
CH_MULT: [ 1, 2, 4, 4 ]
|
||||
Z_CHANNELS: 4
|
||||
DOUBLE_Z: True
|
||||
DROPOUT: 0.0
|
||||
RESAMP_WITH_CONV: True
|
||||
#
|
||||
DECODER:
|
||||
NAME: Decoder
|
||||
CH: 128
|
||||
OUT_CH: 3
|
||||
NUM_RES_BLOCKS: 2
|
||||
IN_CHANNELS: 3
|
||||
ATTN_RESOLUTIONS: [ ]
|
||||
CH_MULT: [ 1, 2, 4, 4 ]
|
||||
Z_CHANNELS: 4
|
||||
DROPOUT: 0.0
|
||||
RESAMP_WITH_CONV: True
|
||||
GIVE_PRE_END: False
|
||||
TANH_OUT: False
|
||||
#
|
||||
COND_STAGE_MODEL:
|
||||
NAME: T5EmbedderHF
|
||||
PRETRAINED_MODEL: ms://AI-ModelScope/PixArt-alpha@t5-v1_1-xxl/
|
||||
TOKENIZER_PATH: ms://AI-ModelScope/PixArt-alpha@t5-v1_1-xxl/
|
||||
LENGTH: 120
|
||||
CLEAN: heavy
|
||||
USE_GRAD: False
|
||||
#
|
||||
LOSS:
|
||||
NAME: ReconstructLoss
|
||||
LOSS_TYPE: l2
|
||||
#
|
||||
SAMPLE_ARGS:
|
||||
SAMPLER: ddim
|
||||
SAMPLE_STEPS: 20
|
||||
SEED: 2024
|
||||
GUIDE_SCALE: 4.5
|
||||
GUIDE_RESCALE: 0.5
|
||||
DISCRETIZATION: trailing
|
||||
RUN_TRAIN_N: False
|
||||
#
|
||||
OPTIMIZER:
|
||||
NAME: AdamW
|
||||
LEARNING_RATE: 0.0001
|
||||
BETAS: [ 0.9, 0.999 ]
|
||||
EPS: 1e-8
|
||||
WEIGHT_DECAY: 1e-2
|
||||
AMSGRAD: False
|
||||
#
|
||||
TRAIN_DATA:
|
||||
NAME: ImageTextPairMSDataset
|
||||
MODE: train
|
||||
MS_DATASET_NAME: style_custom_dataset
|
||||
MS_DATASET_NAMESPACE: damo
|
||||
MS_DATASET_SUBNAME: 3D
|
||||
PROMPT_PREFIX: ""
|
||||
MS_DATASET_SPLIT: train
|
||||
MS_REMAP_KEYS: { 'Image:FILE': 'Target:FILE' }
|
||||
REPLACE_STYLE: False
|
||||
PIN_MEMORY: True
|
||||
BATCH_SIZE: 1
|
||||
NUM_WORKERS: 4
|
||||
SAMPLER:
|
||||
NAME: LoopSampler
|
||||
TRANSFORMS:
|
||||
- NAME: LoadImageFromFile
|
||||
RGB_ORDER: RGB
|
||||
BACKEND: pillow
|
||||
- NAME: FlexibleResize
|
||||
INTERPOLATION: bilinear
|
||||
SIZE: [ 1024, 1024 ]
|
||||
INPUT_KEY: [ 'img' ]
|
||||
OUTPUT_KEY: [ 'img' ]
|
||||
BACKEND: pillow
|
||||
- NAME: FlexibleCenterCrop
|
||||
SIZE: [ 1024, 1024 ]
|
||||
INPUT_KEY: [ 'img' ]
|
||||
OUTPUT_KEY: [ 'img' ]
|
||||
BACKEND: pillow
|
||||
- NAME: ImageToTensor
|
||||
INPUT_KEY: [ 'img' ]
|
||||
OUTPUT_KEY: [ 'img' ]
|
||||
BACKEND: pillow
|
||||
- NAME: Normalize
|
||||
MEAN: [ 0.5, 0.5, 0.5 ]
|
||||
STD: [ 0.5, 0.5, 0.5 ]
|
||||
INPUT_KEY: [ 'img' ]
|
||||
OUTPUT_KEY: [ 'image' ]
|
||||
BACKEND: torchvision
|
||||
- NAME: Select
|
||||
KEYS: [ 'image', 'prompt' ]
|
||||
META_KEYS: [ 'data_key' ]
|
||||
#
|
||||
EVAL_DATA:
|
||||
NAME: Text2ImageDataset
|
||||
MODE: eval
|
||||
PROMPT_FILE:
|
||||
PROMPT_DATA: [ "a boy wearing a jacket", "a dog running on the lawn" ]
|
||||
IMAGE_SIZE: [ 1024, 1024 ]
|
||||
FIELDS: [ "prompt" ]
|
||||
DELIMITER: '#;#'
|
||||
PROMPT_PREFIX: ''
|
||||
PIN_MEMORY: True
|
||||
BATCH_SIZE: 1
|
||||
NUM_WORKERS: 4
|
||||
TRANSFORMS:
|
||||
- NAME: Select
|
||||
KEYS: [ 'index', 'prompt' ]
|
||||
META_KEYS: [ 'image_size' ]
|
||||
#
|
||||
TRAIN_HOOKS:
|
||||
-
|
||||
NAME: BackwardHook
|
||||
PRIORITY: 0
|
||||
-
|
||||
NAME: LogHook
|
||||
LOG_INTERVAL: 10
|
||||
SHOW_GPU_MEM: True
|
||||
-
|
||||
NAME: TensorboardLogHook
|
||||
-
|
||||
NAME: CheckpointHook
|
||||
INTERVAL: 10000
|
||||
PRIORITY: 200
|
||||
SAVE_LAST: True
|
||||
SAVE_NAME_PREFIX: 'step'
|
||||
DISABLE_SNAPSHOT: True
|
||||
#
|
||||
EVAL_HOOKS:
|
||||
-
|
||||
NAME: ProbeDataHook
|
||||
PROB_INTERVAL: 100
|
||||
SAVE_LAST: True
|
||||
SAVE_NAME_PREFIX: 'step'
|
||||
SAVE_PROBE_PREFIX: 'image'
|
||||
@@ -0,0 +1,233 @@
|
||||
ENV:
|
||||
BACKEND: nccl
|
||||
#
|
||||
SOLVER:
|
||||
NAME: LatentDiffusionSolver
|
||||
RESUME_FROM:
|
||||
LOAD_MODEL_ONLY: True
|
||||
USE_FSDP: False
|
||||
SHARDING_STRATEGY:
|
||||
USE_AMP: True
|
||||
DTYPE: float16
|
||||
CHANNELS_LAST: True
|
||||
MAX_STEPS: 500
|
||||
MAX_EPOCHS: -1
|
||||
NUM_FOLDS: 1
|
||||
ACCU_STEP: 1
|
||||
EVAL_INTERVAL: 50
|
||||
RESCALE_LR: False
|
||||
#
|
||||
WORK_DIR: ./cache/save_data/dit_sd3_1024
|
||||
LOG_FILE: std_log.txt
|
||||
#
|
||||
FILE_SYSTEM:
|
||||
- NAME: "ModelscopeFs"
|
||||
TEMP_DIR: ./cache/cache_data
|
||||
#
|
||||
MODEL:
|
||||
NAME: LatentDiffusionSD3
|
||||
PARAMETERIZATION: rf
|
||||
TIMESTEPS: 1000
|
||||
MIN_SNR_GAMMA:
|
||||
ZERO_TERMINAL_SNR: False
|
||||
PRETRAINED_MODEL:
|
||||
IGNORE_KEYS: [ ]
|
||||
SCALE_FACTOR: 1.5305
|
||||
SHIFT_FACTOR: 0.0609
|
||||
DEFAULT_N_PROMPT:
|
||||
SCHEDULE_ARGS:
|
||||
"NAME": "shifted"
|
||||
"SHIFT": 3
|
||||
USE_EMA: False
|
||||
T_WEIGHT: uniform
|
||||
#
|
||||
DIFFUSION_MODEL:
|
||||
NAME: MMDiT
|
||||
PRETRAINED_MODEL: ms://AI-ModelScope/stable-diffusion-3-medium@sd3_medium.safetensors
|
||||
IGNORE_KEYS: '^first_stage_model.'
|
||||
IN_CHANNELS: 16
|
||||
PATCH_SIZE: 2
|
||||
OUT_CHANNELS: 16
|
||||
DEPTH: 24
|
||||
INPUT_SIZE:
|
||||
ADM_IN_CHANNELS: 2048
|
||||
CONTEXT_EMBEDDER_CONFIG: { 'target': 'torch.nn.Linear', 'params': { 'in_features': 4096, 'out_features': 1536 } }
|
||||
NUM_PATCHES: 36864
|
||||
POS_EMBED_MAX_SIZE: 192
|
||||
POS_EMBED_SCALING_FACTOR:
|
||||
USE_CHECKPOINT: True
|
||||
#
|
||||
FIRST_STAGE_MODEL:
|
||||
NAME: AutoencoderKL
|
||||
PRETRAINED_MODEL: ms://AI-ModelScope/stable-diffusion-3-medium@sd3_medium.safetensors
|
||||
EMBED_DIM: 16
|
||||
IGNORE_KEYS: '^model.diffusion_model.'
|
||||
BATCH_SIZE: 1
|
||||
USE_CONV: False
|
||||
#
|
||||
ENCODER:
|
||||
NAME: Encoder
|
||||
CH: 128
|
||||
OUT_CH: 3
|
||||
NUM_RES_BLOCKS: 2
|
||||
IN_CHANNELS: 3
|
||||
ATTN_RESOLUTIONS: [ ]
|
||||
CH_MULT: [ 1, 2, 4, 4 ]
|
||||
Z_CHANNELS: 16
|
||||
DOUBLE_Z: True
|
||||
DROPOUT: 0.0
|
||||
RESAMP_WITH_CONV: True
|
||||
#
|
||||
DECODER:
|
||||
NAME: Decoder
|
||||
CH: 128
|
||||
OUT_CH: 3
|
||||
NUM_RES_BLOCKS: 2
|
||||
IN_CHANNELS: 3
|
||||
ATTN_RESOLUTIONS: [ ]
|
||||
CH_MULT: [ 1, 2, 4, 4 ]
|
||||
Z_CHANNELS: 16
|
||||
DROPOUT: 0.0
|
||||
RESAMP_WITH_CONV: True
|
||||
GIVE_PRE_END: False
|
||||
TANH_OUT: False
|
||||
#
|
||||
COND_STAGE_MODEL:
|
||||
NAME: SD3TextEmbedder
|
||||
P_ZERO: 0.0
|
||||
CLIP_L:
|
||||
NAME: FrozenCLIPEmbedder2
|
||||
PRETRAINED_MODEL: ms://AI-ModelScope/stable-diffusion-3-medium-diffusers@text_encoder
|
||||
TOKENIZER_PATH: ms://AI-ModelScope/stable-diffusion-3-medium-diffusers@tokenizer
|
||||
MAX_LENGTH: 77
|
||||
FREEZE: True
|
||||
LAYER: penultimate
|
||||
RETURN_POOLED: True
|
||||
USE_FINAL_LAYER_NORM: False
|
||||
IS_TRAINABLE: False
|
||||
CLIP_G:
|
||||
NAME: FrozenCLIPEmbedder2
|
||||
PRETRAINED_MODEL: ms://AI-ModelScope/stable-diffusion-3-medium-diffusers@text_encoder_2
|
||||
TOKENIZER_PATH: ms://AI-ModelScope/stable-diffusion-3-medium-diffusers@tokenizer_2
|
||||
MAX_LENGTH: 77
|
||||
FREEZE: True
|
||||
LAYER: penultimate
|
||||
RETURN_POOLED: True
|
||||
USE_FINAL_LAYER_NORM: False
|
||||
IS_TRAINABLE: False
|
||||
T5_XXL:
|
||||
NAME: T5EmbedderHF
|
||||
PRETRAINED_MODEL: ms://AI-ModelScope/stable-diffusion-3-medium-diffusers@text_encoder_3
|
||||
TOKENIZER_PATH: ms://AI-ModelScope/stable-diffusion-3-medium-diffusers@tokenizer_3
|
||||
LENGTH: 256
|
||||
CLEAN: whitespace
|
||||
USE_GRAD: False
|
||||
T5_DTYPE: float16
|
||||
#
|
||||
LOSS:
|
||||
NAME: ReconstructLoss
|
||||
LOSS_TYPE: l2
|
||||
#
|
||||
SAMPLE_ARGS:
|
||||
SAMPLER: euler
|
||||
SAMPLE_STEPS: 28
|
||||
SEED: 1749023094
|
||||
GUIDE_SCALE: 5.0
|
||||
GUIDE_RESCALE: 0.0
|
||||
DISCRETIZATION: trailing
|
||||
RUN_TRAIN_N: False
|
||||
#
|
||||
OPTIMIZER:
|
||||
NAME: AdamW
|
||||
LEARNING_RATE: 1e-5
|
||||
BETAS: [ 0.9, 0.999 ]
|
||||
EPS: 1e-8
|
||||
WEIGHT_DECAY: 1e-2
|
||||
AMSGRAD: False
|
||||
#
|
||||
TRAIN_DATA:
|
||||
NAME: ImageTextPairMSDataset
|
||||
MODE: train
|
||||
MS_DATASET_NAME: style_custom_dataset
|
||||
MS_DATASET_NAMESPACE: damo
|
||||
MS_DATASET_SUBNAME: 3D
|
||||
PROMPT_PREFIX: ""
|
||||
MS_DATASET_SPLIT: train
|
||||
MS_REMAP_KEYS: { 'Image:FILE': 'Target:FILE' }
|
||||
REPLACE_STYLE: False
|
||||
PIN_MEMORY: True
|
||||
BATCH_SIZE: 1
|
||||
NUM_WORKERS: 4
|
||||
SAMPLER:
|
||||
NAME: LoopSampler
|
||||
TRANSFORMS:
|
||||
- NAME: LoadImageFromFile
|
||||
RGB_ORDER: RGB
|
||||
BACKEND: pillow
|
||||
- NAME: FlexibleResize
|
||||
INTERPOLATION: bilinear
|
||||
SIZE: [ 1024, 1024 ]
|
||||
INPUT_KEY: [ 'img' ]
|
||||
OUTPUT_KEY: [ 'img' ]
|
||||
BACKEND: pillow
|
||||
- NAME: FlexibleCenterCrop
|
||||
SIZE: [ 1024, 1024 ]
|
||||
INPUT_KEY: [ 'img' ]
|
||||
OUTPUT_KEY: [ 'img' ]
|
||||
BACKEND: pillow
|
||||
- NAME: ImageToTensor
|
||||
INPUT_KEY: [ 'img' ]
|
||||
OUTPUT_KEY: [ 'img' ]
|
||||
BACKEND: pillow
|
||||
- NAME: Normalize
|
||||
MEAN: [ 0.5, 0.5, 0.5 ]
|
||||
STD: [ 0.5, 0.5, 0.5 ]
|
||||
INPUT_KEY: [ 'img' ]
|
||||
OUTPUT_KEY: [ 'image' ]
|
||||
BACKEND: torchvision
|
||||
- NAME: Select
|
||||
KEYS: [ 'image', 'prompt' ]
|
||||
META_KEYS: [ 'data_key' ]
|
||||
#
|
||||
EVAL_DATA:
|
||||
NAME: Text2ImageDataset
|
||||
MODE: eval
|
||||
PROMPT_FILE:
|
||||
PROMPT_DATA: [ "a cat holds a blackboard that writes \"hello world\"", "a dog running on the lawn" ]
|
||||
IMAGE_SIZE: [ 1024, 1024 ]
|
||||
FIELDS: [ "prompt" ]
|
||||
DELIMITER: '#;#'
|
||||
PROMPT_PREFIX: ''
|
||||
PIN_MEMORY: True
|
||||
BATCH_SIZE: 1
|
||||
NUM_WORKERS: 4
|
||||
TRANSFORMS:
|
||||
- NAME: Select
|
||||
KEYS: [ 'index', 'prompt' ]
|
||||
META_KEYS: [ 'image_size' ]
|
||||
#
|
||||
TRAIN_HOOKS:
|
||||
-
|
||||
NAME: BackwardHook
|
||||
PRIORITY: 10000
|
||||
-
|
||||
NAME: LogHook
|
||||
LOG_INTERVAL: 10
|
||||
SHOW_GPU_MEM: True
|
||||
-
|
||||
NAME: TensorboardLogHook
|
||||
-
|
||||
NAME: CheckpointHook
|
||||
INTERVAL: 10000
|
||||
PRIORITY: 200
|
||||
SAVE_LAST: True
|
||||
SAVE_NAME_PREFIX: 'step'
|
||||
DISABLE_SNAPSHOT: True
|
||||
#
|
||||
EVAL_HOOKS:
|
||||
-
|
||||
NAME: ProbeDataHook
|
||||
PROB_INTERVAL: 50
|
||||
SAVE_LAST: True
|
||||
SAVE_NAME_PREFIX: 'step'
|
||||
SAVE_PROBE_PREFIX: 'image'
|
||||
@@ -0,0 +1,241 @@
|
||||
ENV:
|
||||
BACKEND: nccl
|
||||
#
|
||||
SOLVER:
|
||||
NAME: LatentDiffusionSolver
|
||||
RESUME_FROM:
|
||||
LOAD_MODEL_ONLY: True
|
||||
USE_FSDP: False
|
||||
SHARDING_STRATEGY:
|
||||
USE_AMP: True
|
||||
DTYPE: float16
|
||||
CHANNELS_LAST: True
|
||||
MAX_STEPS: 500
|
||||
MAX_EPOCHS: -1
|
||||
NUM_FOLDS: 1
|
||||
ACCU_STEP: 1
|
||||
EVAL_INTERVAL: 50
|
||||
RESCALE_LR: False
|
||||
#
|
||||
WORK_DIR: ./cache/save_data/dit_sd3_1024_lora
|
||||
LOG_FILE: std_log.txt
|
||||
#
|
||||
FILE_SYSTEM:
|
||||
- NAME: "ModelscopeFs"
|
||||
TEMP_DIR: ./cache/cache_data
|
||||
#
|
||||
TUNER:
|
||||
- NAME: SwiftLoRA
|
||||
R: 128
|
||||
LORA_ALPHA: 128
|
||||
LORA_DROPOUT: 0.0
|
||||
BIAS: "none"
|
||||
TARGET_MODULES: "model.*(.attn.qkv|.attn.proj|mlp.fc1|mlp.fc2)$"
|
||||
#
|
||||
MODEL:
|
||||
NAME: LatentDiffusionSD3
|
||||
PARAMETERIZATION: rf
|
||||
TIMESTEPS: 1000
|
||||
MIN_SNR_GAMMA:
|
||||
ZERO_TERMINAL_SNR: False
|
||||
PRETRAINED_MODEL:
|
||||
IGNORE_KEYS: [ ]
|
||||
SCALE_FACTOR: 1.5305
|
||||
SHIFT_FACTOR: 0.0609
|
||||
DEFAULT_N_PROMPT:
|
||||
SCHEDULE_ARGS:
|
||||
"NAME": "shifted"
|
||||
"SHIFT": 3
|
||||
USE_EMA: False
|
||||
T_WEIGHT: uniform
|
||||
#
|
||||
DIFFUSION_MODEL:
|
||||
NAME: MMDiT
|
||||
PRETRAINED_MODEL: ms://AI-ModelScope/stable-diffusion-3-medium@sd3_medium.safetensors
|
||||
IGNORE_KEYS: '^first_stage_model.'
|
||||
IN_CHANNELS: 16
|
||||
PATCH_SIZE: 2
|
||||
OUT_CHANNELS: 16
|
||||
DEPTH: 24
|
||||
INPUT_SIZE:
|
||||
ADM_IN_CHANNELS: 2048
|
||||
CONTEXT_EMBEDDER_CONFIG: { 'target': 'torch.nn.Linear', 'params': { 'in_features': 4096, 'out_features': 1536 } }
|
||||
NUM_PATCHES: 36864
|
||||
POS_EMBED_MAX_SIZE: 192
|
||||
POS_EMBED_SCALING_FACTOR:
|
||||
USE_CHECKPOINT: True
|
||||
#
|
||||
FIRST_STAGE_MODEL:
|
||||
NAME: AutoencoderKL
|
||||
PRETRAINED_MODEL: ms://AI-ModelScope/stable-diffusion-3-medium@sd3_medium.safetensors
|
||||
EMBED_DIM: 16
|
||||
IGNORE_KEYS: '^model.diffusion_model.'
|
||||
BATCH_SIZE: 1
|
||||
USE_CONV: False
|
||||
#
|
||||
ENCODER:
|
||||
NAME: Encoder
|
||||
CH: 128
|
||||
OUT_CH: 3
|
||||
NUM_RES_BLOCKS: 2
|
||||
IN_CHANNELS: 3
|
||||
ATTN_RESOLUTIONS: [ ]
|
||||
CH_MULT: [ 1, 2, 4, 4 ]
|
||||
Z_CHANNELS: 16
|
||||
DOUBLE_Z: True
|
||||
DROPOUT: 0.0
|
||||
RESAMP_WITH_CONV: True
|
||||
#
|
||||
DECODER:
|
||||
NAME: Decoder
|
||||
CH: 128
|
||||
OUT_CH: 3
|
||||
NUM_RES_BLOCKS: 2
|
||||
IN_CHANNELS: 3
|
||||
ATTN_RESOLUTIONS: [ ]
|
||||
CH_MULT: [ 1, 2, 4, 4 ]
|
||||
Z_CHANNELS: 16
|
||||
DROPOUT: 0.0
|
||||
RESAMP_WITH_CONV: True
|
||||
GIVE_PRE_END: False
|
||||
TANH_OUT: False
|
||||
#
|
||||
COND_STAGE_MODEL:
|
||||
NAME: SD3TextEmbedder
|
||||
P_ZERO: 0.0
|
||||
CLIP_L:
|
||||
NAME: FrozenCLIPEmbedder2
|
||||
PRETRAINED_MODEL: ms://AI-ModelScope/stable-diffusion-3-medium-diffusers@text_encoder
|
||||
TOKENIZER_PATH: ms://AI-ModelScope/stable-diffusion-3-medium-diffusers@tokenizer
|
||||
MAX_LENGTH: 77
|
||||
FREEZE: True
|
||||
LAYER: penultimate
|
||||
RETURN_POOLED: True
|
||||
USE_FINAL_LAYER_NORM: False
|
||||
IS_TRAINABLE: False
|
||||
CLIP_G:
|
||||
NAME: FrozenCLIPEmbedder2
|
||||
PRETRAINED_MODEL: ms://AI-ModelScope/stable-diffusion-3-medium-diffusers@text_encoder_2
|
||||
TOKENIZER_PATH: ms://AI-ModelScope/stable-diffusion-3-medium-diffusers@tokenizer_2
|
||||
MAX_LENGTH: 77
|
||||
FREEZE: True
|
||||
LAYER: penultimate
|
||||
RETURN_POOLED: True
|
||||
USE_FINAL_LAYER_NORM: False
|
||||
IS_TRAINABLE: False
|
||||
T5_XXL:
|
||||
NAME: T5EmbedderHF
|
||||
PRETRAINED_MODEL: ms://AI-ModelScope/stable-diffusion-3-medium-diffusers@text_encoder_3
|
||||
TOKENIZER_PATH: ms://AI-ModelScope/stable-diffusion-3-medium-diffusers@tokenizer_3
|
||||
LENGTH: 256
|
||||
CLEAN: whitespace
|
||||
USE_GRAD: False
|
||||
T5_DTYPE: float16
|
||||
#
|
||||
LOSS:
|
||||
NAME: ReconstructLoss
|
||||
LOSS_TYPE: l2
|
||||
#
|
||||
SAMPLE_ARGS:
|
||||
SAMPLER: euler
|
||||
SAMPLE_STEPS: 28
|
||||
SEED: 1749023094
|
||||
GUIDE_SCALE: 5.0
|
||||
GUIDE_RESCALE: 0.0
|
||||
DISCRETIZATION: trailing
|
||||
RUN_TRAIN_N: False
|
||||
#
|
||||
OPTIMIZER:
|
||||
NAME: AdamW
|
||||
LEARNING_RATE: 5e-5
|
||||
BETAS: [ 0.9, 0.999 ]
|
||||
EPS: 1e-8
|
||||
WEIGHT_DECAY: 1e-2
|
||||
AMSGRAD: False
|
||||
#
|
||||
TRAIN_DATA:
|
||||
NAME: ImageTextPairMSDataset
|
||||
MODE: train
|
||||
MS_DATASET_NAME: style_custom_dataset
|
||||
MS_DATASET_NAMESPACE: damo
|
||||
MS_DATASET_SUBNAME: 3D
|
||||
PROMPT_PREFIX: ""
|
||||
MS_DATASET_SPLIT: train
|
||||
MS_REMAP_KEYS: { 'Image:FILE': 'Target:FILE' }
|
||||
REPLACE_STYLE: False
|
||||
PIN_MEMORY: True
|
||||
BATCH_SIZE: 1
|
||||
NUM_WORKERS: 4
|
||||
SAMPLER:
|
||||
NAME: LoopSampler
|
||||
TRANSFORMS:
|
||||
- NAME: LoadImageFromFile
|
||||
RGB_ORDER: RGB
|
||||
BACKEND: pillow
|
||||
- NAME: FlexibleResize
|
||||
INTERPOLATION: bilinear
|
||||
SIZE: [ 1024, 1024 ]
|
||||
INPUT_KEY: [ 'img' ]
|
||||
OUTPUT_KEY: [ 'img' ]
|
||||
BACKEND: pillow
|
||||
- NAME: FlexibleCenterCrop
|
||||
SIZE: [ 1024, 1024 ]
|
||||
INPUT_KEY: [ 'img' ]
|
||||
OUTPUT_KEY: [ 'img' ]
|
||||
BACKEND: pillow
|
||||
- NAME: ImageToTensor
|
||||
INPUT_KEY: [ 'img' ]
|
||||
OUTPUT_KEY: [ 'img' ]
|
||||
BACKEND: pillow
|
||||
- NAME: Normalize
|
||||
MEAN: [ 0.5, 0.5, 0.5 ]
|
||||
STD: [ 0.5, 0.5, 0.5 ]
|
||||
INPUT_KEY: [ 'img' ]
|
||||
OUTPUT_KEY: [ 'image' ]
|
||||
BACKEND: torchvision
|
||||
- NAME: Select
|
||||
KEYS: [ 'image', 'prompt' ]
|
||||
META_KEYS: [ 'data_key' ]
|
||||
#
|
||||
EVAL_DATA:
|
||||
NAME: Text2ImageDataset
|
||||
MODE: eval
|
||||
PROMPT_FILE:
|
||||
PROMPT_DATA: [ "a cat holds a blackboard that writes \"hello world\"", "a dog running on the lawn" ]
|
||||
IMAGE_SIZE: [ 1024, 1024 ]
|
||||
FIELDS: [ "prompt" ]
|
||||
DELIMITER: '#;#'
|
||||
PROMPT_PREFIX: ''
|
||||
PIN_MEMORY: True
|
||||
BATCH_SIZE: 1
|
||||
NUM_WORKERS: 4
|
||||
TRANSFORMS:
|
||||
- NAME: Select
|
||||
KEYS: [ 'index', 'prompt' ]
|
||||
META_KEYS: [ 'image_size' ]
|
||||
#
|
||||
TRAIN_HOOKS:
|
||||
-
|
||||
NAME: BackwardHook
|
||||
PRIORITY: 10000
|
||||
-
|
||||
NAME: LogHook
|
||||
LOG_INTERVAL: 10
|
||||
SHOW_GPU_MEM: True
|
||||
-
|
||||
NAME: TensorboardLogHook
|
||||
-
|
||||
NAME: CheckpointHook
|
||||
INTERVAL: 10000
|
||||
PRIORITY: 200
|
||||
SAVE_LAST: True
|
||||
SAVE_NAME_PREFIX: 'step'
|
||||
DISABLE_SNAPSHOT: True
|
||||
#
|
||||
EVAL_HOOKS:
|
||||
-
|
||||
NAME: ProbeDataHook
|
||||
PROB_INTERVAL: 50
|
||||
SAVE_LAST: True
|
||||
SAVE_NAME_PREFIX: 'step'
|
||||
SAVE_PROBE_PREFIX: 'image'
|
||||
@@ -14,13 +14,14 @@ SOLVER:
|
||||
NUM_FOLDS: 1
|
||||
ACCU_STEP: 1
|
||||
EVAL_INTERVAL: 100
|
||||
RESCALE_LR: False
|
||||
#
|
||||
WORK_DIR: ./cache/save_data/sd15_512_full
|
||||
LOG_FILE: std_log.txt
|
||||
#
|
||||
FILE_SYSTEM:
|
||||
NAME: "ModelscopeFs"
|
||||
TEMP_DIR: "./cache/data"
|
||||
TEMP_DIR: "./cache/cache_data"
|
||||
#
|
||||
MODEL:
|
||||
NAME: LatentDiffusion
|
||||
@@ -124,7 +125,7 @@ SOLVER:
|
||||
#
|
||||
OPTIMIZER:
|
||||
NAME: AdamW
|
||||
LEARNING_RATE: 0.0064
|
||||
LEARNING_RATE: 0.00001
|
||||
BETAS: [ 0.9, 0.999 ]
|
||||
EPS: 1e-8
|
||||
WEIGHT_DECAY: 1e-2
|
||||
@@ -190,7 +191,7 @@ SOLVER:
|
||||
NUM_WORKERS: 4
|
||||
FILE_SYSTEM:
|
||||
NAME: "ModelscopeFs"
|
||||
TEMP_DIR: "./cache/data"
|
||||
TEMP_DIR: "./cache/cache_data"
|
||||
#
|
||||
TRANSFORMS:
|
||||
-
|
||||
|
||||
@@ -14,13 +14,14 @@ SOLVER:
|
||||
NUM_FOLDS: 1
|
||||
ACCU_STEP: 1
|
||||
EVAL_INTERVAL: 100
|
||||
RESCALE_LR: False
|
||||
#
|
||||
WORK_DIR: ./cache/save_data/sd15_512_lora
|
||||
LOG_FILE: std_log.txt
|
||||
#
|
||||
FILE_SYSTEM:
|
||||
NAME: "ModelscopeFs"
|
||||
TEMP_DIR: "./cache/data"
|
||||
TEMP_DIR: "./cache/cache_data"
|
||||
TUNER:
|
||||
-
|
||||
NAME: SwiftLoRA
|
||||
@@ -132,7 +133,7 @@ SOLVER:
|
||||
#
|
||||
OPTIMIZER:
|
||||
NAME: AdamW
|
||||
LEARNING_RATE: 0.064
|
||||
LEARNING_RATE: 0.0001
|
||||
BETAS: [ 0.9, 0.999 ]
|
||||
EPS: 1e-8
|
||||
WEIGHT_DECAY: 1e-2
|
||||
@@ -198,7 +199,7 @@ SOLVER:
|
||||
NUM_WORKERS: 4
|
||||
FILE_SYSTEM:
|
||||
NAME: "ModelscopeFs"
|
||||
TEMP_DIR: "./cache/data"
|
||||
TEMP_DIR: "./cache/cache_data"
|
||||
#
|
||||
TRANSFORMS:
|
||||
-
|
||||
|
||||
@@ -14,13 +14,14 @@ SOLVER:
|
||||
NUM_FOLDS: 1
|
||||
ACCU_STEP: 1
|
||||
EVAL_INTERVAL: 100
|
||||
RESCALE_LR: False
|
||||
#
|
||||
WORK_DIR: ./cache/save_data/sd15_512_textlora
|
||||
LOG_FILE: std_log.txt
|
||||
#
|
||||
FILE_SYSTEM:
|
||||
NAME: "ModelscopeFs"
|
||||
TEMP_DIR: "./cache/data"
|
||||
TEMP_DIR: "./cache/cache_data"
|
||||
TUNER:
|
||||
-
|
||||
NAME: SwiftLoRA
|
||||
@@ -140,7 +141,7 @@ SOLVER:
|
||||
#
|
||||
OPTIMIZER:
|
||||
NAME: AdamW
|
||||
LEARNING_RATE: 0.064
|
||||
LEARNING_RATE: 0.0001
|
||||
BETAS: [ 0.9, 0.999 ]
|
||||
EPS: 1e-8
|
||||
WEIGHT_DECAY: 1e-2
|
||||
@@ -206,7 +207,7 @@ SOLVER:
|
||||
NUM_WORKERS: 4
|
||||
FILE_SYSTEM:
|
||||
NAME: "ModelscopeFs"
|
||||
TEMP_DIR: "./cache/data"
|
||||
TEMP_DIR: "./cache/cache_data"
|
||||
#
|
||||
TRANSFORMS:
|
||||
-
|
||||
|
||||
@@ -14,13 +14,14 @@ SOLVER:
|
||||
NUM_FOLDS: 1
|
||||
ACCU_STEP: 1
|
||||
EVAL_INTERVAL: 100
|
||||
RESCALE_LR: False
|
||||
#
|
||||
WORK_DIR: ./cache/save_data/sd21_512_full
|
||||
LOG_FILE: std_log.txt
|
||||
#
|
||||
FILE_SYSTEM:
|
||||
NAME: "ModelscopeFs"
|
||||
TEMP_DIR: "./cache/data"
|
||||
TEMP_DIR: "./cache/cache_data"
|
||||
#
|
||||
MODEL:
|
||||
NAME: LatentDiffusion
|
||||
@@ -120,7 +121,7 @@ SOLVER:
|
||||
#
|
||||
OPTIMIZER:
|
||||
NAME: AdamW
|
||||
LEARNING_RATE: 0.0064
|
||||
LEARNING_RATE: 0.00001
|
||||
BETAS: [ 0.9, 0.999 ]
|
||||
EPS: 1e-8
|
||||
WEIGHT_DECAY: 1e-2
|
||||
@@ -186,7 +187,7 @@ SOLVER:
|
||||
NUM_WORKERS: 4
|
||||
FILE_SYSTEM:
|
||||
NAME: "ModelscopeFs"
|
||||
TEMP_DIR: "./cache/data"
|
||||
TEMP_DIR: "./cache/cache_data"
|
||||
#
|
||||
TRANSFORMS:
|
||||
-
|
||||
|
||||
@@ -14,13 +14,14 @@ SOLVER:
|
||||
NUM_FOLDS: 1
|
||||
ACCU_STEP: 1
|
||||
EVAL_INTERVAL: 100
|
||||
RESCALE_LR: False
|
||||
#
|
||||
WORK_DIR: ./cache/save_data/sd21_512_lora
|
||||
LOG_FILE: std_log.txt
|
||||
#
|
||||
FILE_SYSTEM:
|
||||
NAME: "ModelscopeFs"
|
||||
TEMP_DIR: "./cache/data"
|
||||
TEMP_DIR: "./cache/cache_data"
|
||||
#
|
||||
TUNER:
|
||||
-
|
||||
@@ -129,7 +130,7 @@ SOLVER:
|
||||
#
|
||||
OPTIMIZER:
|
||||
NAME: AdamW
|
||||
LEARNING_RATE: 0.0064
|
||||
LEARNING_RATE: 0.00001
|
||||
BETAS: [ 0.9, 0.999 ]
|
||||
EPS: 1e-8
|
||||
WEIGHT_DECAY: 1e-2
|
||||
@@ -195,7 +196,7 @@ SOLVER:
|
||||
NUM_WORKERS: 4
|
||||
FILE_SYSTEM:
|
||||
NAME: "ModelscopeFs"
|
||||
TEMP_DIR: "./cache/data"
|
||||
TEMP_DIR: "./cache/cache_data"
|
||||
#
|
||||
TRANSFORMS:
|
||||
-
|
||||
|
||||
@@ -14,13 +14,14 @@ SOLVER:
|
||||
NUM_FOLDS: 1
|
||||
ACCU_STEP: 1
|
||||
EVAL_INTERVAL: 100
|
||||
RESCALE_LR: False
|
||||
#
|
||||
WORK_DIR: ./cache/save_data/sd21_768_full
|
||||
LOG_FILE: std_log.txt
|
||||
#
|
||||
FILE_SYSTEM:
|
||||
NAME: "ModelscopeFs"
|
||||
TEMP_DIR: "./cache/data"
|
||||
TEMP_DIR: "./cache/cache_data"
|
||||
#
|
||||
MODEL:
|
||||
NAME: LatentDiffusion
|
||||
@@ -120,7 +121,7 @@ SOLVER:
|
||||
#
|
||||
OPTIMIZER:
|
||||
NAME: AdamW
|
||||
LEARNING_RATE: 0.0064
|
||||
LEARNING_RATE: 0.00001
|
||||
BETAS: [ 0.9, 0.999 ]
|
||||
EPS: 1e-8
|
||||
WEIGHT_DECAY: 1e-2
|
||||
@@ -186,7 +187,7 @@ SOLVER:
|
||||
NUM_WORKERS: 4
|
||||
FILE_SYSTEM:
|
||||
NAME: "ModelscopeFs"
|
||||
TEMP_DIR: "./cache/data"
|
||||
TEMP_DIR: "./cache/cache_data"
|
||||
#
|
||||
TRANSFORMS:
|
||||
-
|
||||
|
||||
@@ -14,13 +14,14 @@ SOLVER:
|
||||
NUM_FOLDS: 1
|
||||
ACCU_STEP: 1
|
||||
EVAL_INTERVAL: 100
|
||||
RESCALE_LR: False
|
||||
#
|
||||
WORK_DIR: ./cache/save_data/sd21_768_lora
|
||||
LOG_FILE: std_log.txt
|
||||
#
|
||||
FILE_SYSTEM:
|
||||
NAME: "ModelscopeFs"
|
||||
TEMP_DIR: "./cache/data"
|
||||
TEMP_DIR: "./cache/cache_data"
|
||||
#
|
||||
TUNER:
|
||||
-
|
||||
@@ -129,7 +130,7 @@ SOLVER:
|
||||
#
|
||||
OPTIMIZER:
|
||||
NAME: AdamW
|
||||
LEARNING_RATE: 0.064
|
||||
LEARNING_RATE: 0.0001
|
||||
BETAS: [ 0.9, 0.999 ]
|
||||
EPS: 1e-8
|
||||
WEIGHT_DECAY: 1e-2
|
||||
@@ -195,7 +196,7 @@ SOLVER:
|
||||
NUM_WORKERS: 4
|
||||
FILE_SYSTEM:
|
||||
NAME: "ModelscopeFs"
|
||||
TEMP_DIR: "./cache/data"
|
||||
TEMP_DIR: "./cache/cache_data"
|
||||
#
|
||||
TRANSFORMS:
|
||||
-
|
||||
|
||||
@@ -14,13 +14,14 @@ SOLVER:
|
||||
NUM_FOLDS: 1
|
||||
ACCU_STEP: 1
|
||||
EVAL_INTERVAL: 100
|
||||
RESCALE_LR: False
|
||||
#
|
||||
WORK_DIR: ./cache/save_data/sdxl_1024_full
|
||||
LOG_FILE: std_log.txt
|
||||
#
|
||||
FILE_SYSTEM:
|
||||
NAME: "ModelscopeFs"
|
||||
TEMP_DIR: "./cache/data"
|
||||
TEMP_DIR: "./cache/cache_data"
|
||||
#
|
||||
MODEL:
|
||||
NAME: LatentDiffusionXL
|
||||
@@ -238,7 +239,7 @@ SOLVER:
|
||||
#
|
||||
OPTIMIZER:
|
||||
NAME: AdamW
|
||||
LEARNING_RATE: 0.0064
|
||||
LEARNING_RATE: 0.00001
|
||||
BETAS: [ 0.9, 0.999 ]
|
||||
EPS: 1e-8
|
||||
WEIGHT_DECAY: 1e-2
|
||||
@@ -306,7 +307,7 @@ SOLVER:
|
||||
NUM_WORKERS: 4
|
||||
FILE_SYSTEM:
|
||||
NAME: "ModelscopeFs"
|
||||
TEMP_DIR: "./cache/data"
|
||||
TEMP_DIR: "./cache/cache_data"
|
||||
#
|
||||
TRANSFORMS:
|
||||
- NAME: Select
|
||||
|
||||
@@ -14,13 +14,14 @@ SOLVER:
|
||||
NUM_FOLDS: 1
|
||||
ACCU_STEP: 1
|
||||
EVAL_INTERVAL: 100
|
||||
RESCALE_LR: False
|
||||
#
|
||||
WORK_DIR: ./cache/save_data/sdxl_1024_lora
|
||||
LOG_FILE: std_log.txt
|
||||
#
|
||||
FILE_SYSTEM:
|
||||
NAME: "ModelscopeFs"
|
||||
TEMP_DIR: "./cache/data"
|
||||
TEMP_DIR: "./cache/cache_data"
|
||||
#
|
||||
TUNER:
|
||||
-
|
||||
@@ -247,7 +248,7 @@ SOLVER:
|
||||
#
|
||||
OPTIMIZER:
|
||||
NAME: AdamW
|
||||
LEARNING_RATE: 0.064
|
||||
LEARNING_RATE: 0.0001
|
||||
BETAS: [ 0.9, 0.999 ]
|
||||
EPS: 1e-8
|
||||
WEIGHT_DECAY: 1e-2
|
||||
@@ -315,7 +316,7 @@ SOLVER:
|
||||
NUM_WORKERS: 4
|
||||
FILE_SYSTEM:
|
||||
NAME: "ModelscopeFs"
|
||||
TEMP_DIR: "./cache/data"
|
||||
TEMP_DIR: "./cache/cache_data"
|
||||
#
|
||||
TRANSFORMS:
|
||||
- NAME: Select
|
||||
|
||||
@@ -14,13 +14,14 @@ SOLVER:
|
||||
NUM_FOLDS: 1
|
||||
ACCU_STEP: 1
|
||||
EVAL_INTERVAL: 100
|
||||
RESCALE_LR: False
|
||||
#
|
||||
WORK_DIR: ./cache/save_data/sdxl_1024_textlora
|
||||
LOG_FILE: std_log.txt
|
||||
#
|
||||
FILE_SYSTEM:
|
||||
NAME: "ModelscopeFs"
|
||||
TEMP_DIR: "./cache/data"
|
||||
TEMP_DIR: "./cache/cache_data"
|
||||
#
|
||||
TUNER:
|
||||
-
|
||||
@@ -255,7 +256,7 @@ SOLVER:
|
||||
#
|
||||
OPTIMIZER:
|
||||
NAME: AdamW
|
||||
LEARNING_RATE: 0.064
|
||||
LEARNING_RATE: 0.0001
|
||||
BETAS: [ 0.9, 0.999 ]
|
||||
EPS: 1e-8
|
||||
WEIGHT_DECAY: 1e-2
|
||||
@@ -323,7 +324,7 @@ SOLVER:
|
||||
NUM_WORKERS: 4
|
||||
FILE_SYSTEM:
|
||||
NAME: "ModelscopeFs"
|
||||
TEMP_DIR: "./cache/data"
|
||||
TEMP_DIR: "./cache/cache_data"
|
||||
#
|
||||
TRANSFORMS:
|
||||
- NAME: Select
|
||||
|
||||
@@ -14,13 +14,14 @@ SOLVER:
|
||||
NUM_FOLDS: 1
|
||||
ACCU_STEP: 1
|
||||
EVAL_INTERVAL: 100
|
||||
RESCALE_LR: False
|
||||
#
|
||||
WORK_DIR: ./cache/save_data/sd15_512_sce_ctr_hed
|
||||
LOG_FILE: std_log.txt
|
||||
#
|
||||
FILE_SYSTEM:
|
||||
NAME: "ModelscopeFs"
|
||||
TEMP_DIR: "./cache/data"
|
||||
TEMP_DIR: "./cache/cache_data"
|
||||
#
|
||||
FREEZE:
|
||||
FREEZE_PART: [ "first_stage_model", "cond_stage_model", "model" ]
|
||||
@@ -141,7 +142,7 @@ SOLVER:
|
||||
#
|
||||
OPTIMIZER:
|
||||
NAME: AdamW
|
||||
LEARNING_RATE: 0.064
|
||||
LEARNING_RATE: 0.0001
|
||||
BETAS: [ 0.9, 0.999 ]
|
||||
EPS: 1e-8
|
||||
WEIGHT_DECAY: 1e-2
|
||||
|
||||
@@ -14,13 +14,14 @@ SOLVER:
|
||||
NUM_FOLDS: 1
|
||||
ACCU_STEP: 1
|
||||
EVAL_INTERVAL: 100
|
||||
RESCALE_LR: False
|
||||
#
|
||||
WORK_DIR: ./cache/save_data/sd21_768_sce_ctr_canny
|
||||
LOG_FILE: std_log.txt
|
||||
#
|
||||
FILE_SYSTEM:
|
||||
NAME: "ModelscopeFs"
|
||||
TEMP_DIR: "./cache/data"
|
||||
TEMP_DIR: "./cache/cache_data"
|
||||
#
|
||||
FREEZE:
|
||||
FREEZE_PART: [ "first_stage_model", "cond_stage_model", "model" ]
|
||||
@@ -140,7 +141,7 @@ SOLVER:
|
||||
#
|
||||
OPTIMIZER:
|
||||
NAME: AdamW
|
||||
LEARNING_RATE: 0.064
|
||||
LEARNING_RATE: 0.0001
|
||||
BETAS: [ 0.9, 0.999 ]
|
||||
EPS: 1e-8
|
||||
WEIGHT_DECAY: 1e-2
|
||||
|
||||
@@ -14,13 +14,14 @@ SOLVER:
|
||||
NUM_FOLDS: 1
|
||||
ACCU_STEP: 1
|
||||
EVAL_INTERVAL: 100
|
||||
RESCALE_LR: False
|
||||
#
|
||||
WORK_DIR: ./cache/save_data/sd21_768_sce_ctr_pose
|
||||
LOG_FILE: std_log.txt
|
||||
#
|
||||
FILE_SYSTEM:
|
||||
NAME: "ModelscopeFs"
|
||||
TEMP_DIR: "./cache/data"
|
||||
TEMP_DIR: "./cache/cache_data"
|
||||
#
|
||||
FREEZE:
|
||||
FREEZE_PART: [ "first_stage_model", "cond_stage_model", "model" ]
|
||||
@@ -140,7 +141,7 @@ SOLVER:
|
||||
#
|
||||
OPTIMIZER:
|
||||
NAME: AdamW
|
||||
LEARNING_RATE: 0.064
|
||||
LEARNING_RATE: 0.0001
|
||||
BETAS: [ 0.9, 0.999 ]
|
||||
EPS: 1e-8
|
||||
WEIGHT_DECAY: 1e-2
|
||||
|
||||
@@ -14,13 +14,14 @@ SOLVER:
|
||||
NUM_FOLDS: 1
|
||||
ACCU_STEP: 1
|
||||
EVAL_INTERVAL: 100
|
||||
RESCALE_LR: False
|
||||
#
|
||||
WORK_DIR: ./cache/save_data/sdxl_1024_sce_ctr_canny
|
||||
LOG_FILE: std_log.txt
|
||||
#
|
||||
FILE_SYSTEM:
|
||||
NAME: "ModelscopeFs"
|
||||
TEMP_DIR: "./cache/data"
|
||||
TEMP_DIR: "./cache/cache_data"
|
||||
#
|
||||
FREEZE:
|
||||
FREEZE_PART: [ "first_stage_model", "cond_stage_model", "model" ]
|
||||
@@ -254,7 +255,7 @@ SOLVER:
|
||||
#
|
||||
OPTIMIZER:
|
||||
NAME: AdamW
|
||||
LEARNING_RATE: 0.064
|
||||
LEARNING_RATE: 0.0001
|
||||
BETAS: [ 0.9, 0.999 ]
|
||||
EPS: 1e-8
|
||||
WEIGHT_DECAY: 1e-2
|
||||
|
||||
@@ -14,13 +14,14 @@ SOLVER:
|
||||
NUM_FOLDS: 1
|
||||
ACCU_STEP: 1
|
||||
EVAL_INTERVAL: 100
|
||||
RESCALE_LR: False
|
||||
#
|
||||
WORK_DIR: ./cache/save_data/sdxl_1024_sce_ctr_color
|
||||
LOG_FILE: std_log.txt
|
||||
#
|
||||
FILE_SYSTEM:
|
||||
NAME: "ModelscopeFs"
|
||||
TEMP_DIR: "./cache/data"
|
||||
TEMP_DIR: "./cache/cache_data"
|
||||
#
|
||||
FREEZE:
|
||||
FREEZE_PART: [ "first_stage_model", "cond_stage_model", "model" ]
|
||||
@@ -255,7 +256,7 @@ SOLVER:
|
||||
#
|
||||
OPTIMIZER:
|
||||
NAME: AdamW
|
||||
LEARNING_RATE: 0.064
|
||||
LEARNING_RATE: 0.0001
|
||||
BETAS: [ 0.9, 0.999 ]
|
||||
EPS: 1e-8
|
||||
WEIGHT_DECAY: 1e-2
|
||||
|
||||
@@ -14,13 +14,14 @@ SOLVER:
|
||||
NUM_FOLDS: 1
|
||||
ACCU_STEP: 1
|
||||
EVAL_INTERVAL: 100
|
||||
RESCALE_LR: False
|
||||
#
|
||||
WORK_DIR: ./cache/save_data/sdxl_1024_sce_ctr_color_datatxt
|
||||
LOG_FILE: std_log.txt
|
||||
#
|
||||
FILE_SYSTEM:
|
||||
NAME: "ModelscopeFs"
|
||||
TEMP_DIR: "./cache/data"
|
||||
TEMP_DIR: "./cache/cache_data"
|
||||
#
|
||||
FREEZE:
|
||||
FREEZE_PART: [ "first_stage_model", "cond_stage_model", "model" ]
|
||||
@@ -255,7 +256,7 @@ SOLVER:
|
||||
#
|
||||
OPTIMIZER:
|
||||
NAME: AdamW
|
||||
LEARNING_RATE: 0.064
|
||||
LEARNING_RATE: 0.0001
|
||||
BETAS: [ 0.9, 0.999 ]
|
||||
EPS: 1e-8
|
||||
WEIGHT_DECAY: 1e-2
|
||||
|
||||
@@ -14,13 +14,14 @@ SOLVER:
|
||||
NUM_FOLDS: 1
|
||||
ACCU_STEP: 1
|
||||
EVAL_INTERVAL: 100
|
||||
RESCALE_LR: False
|
||||
#
|
||||
WORK_DIR: ./cache/save_data/sdxl_1024_sce_ctr_depth
|
||||
LOG_FILE: std_log.txt
|
||||
#
|
||||
FILE_SYSTEM:
|
||||
NAME: "ModelscopeFs"
|
||||
TEMP_DIR: "./cache/data"
|
||||
TEMP_DIR: "./cache/cache_data"
|
||||
#
|
||||
FREEZE:
|
||||
FREEZE_PART: [ "first_stage_model", "cond_stage_model", "model" ]
|
||||
@@ -255,7 +256,7 @@ SOLVER:
|
||||
#
|
||||
OPTIMIZER:
|
||||
NAME: AdamW
|
||||
LEARNING_RATE: 0.064
|
||||
LEARNING_RATE: 0.0001
|
||||
BETAS: [ 0.9, 0.999 ]
|
||||
EPS: 1e-8
|
||||
WEIGHT_DECAY: 1e-2
|
||||
|
||||
@@ -14,13 +14,14 @@ SOLVER:
|
||||
NUM_FOLDS: 1
|
||||
ACCU_STEP: 1
|
||||
EVAL_INTERVAL: 100
|
||||
RESCALE_LR: False
|
||||
#
|
||||
WORK_DIR: ./cache/save_data/sd15_512_sce_t2i
|
||||
LOG_FILE: std_log.txt
|
||||
#
|
||||
FILE_SYSTEM:
|
||||
NAME: "ModelscopeFs"
|
||||
TEMP_DIR: "./cache/data"
|
||||
TEMP_DIR: "./cache/cache_data"
|
||||
#
|
||||
FREEZE:
|
||||
FREEZE_PART: [ "first_stage_model", "cond_stage_model", "model" ]
|
||||
@@ -134,7 +135,7 @@ SOLVER:
|
||||
#
|
||||
OPTIMIZER:
|
||||
NAME: AdamW
|
||||
LEARNING_RATE: 0.064
|
||||
LEARNING_RATE: 0.0001
|
||||
BETAS: [ 0.9, 0.999 ]
|
||||
EPS: 1e-8
|
||||
WEIGHT_DECAY: 1e-2
|
||||
@@ -200,7 +201,7 @@ SOLVER:
|
||||
NUM_WORKERS: 4
|
||||
FILE_SYSTEM:
|
||||
NAME: "ModelscopeFs"
|
||||
TEMP_DIR: "./cache/data"
|
||||
TEMP_DIR: "./cache/cache_data"
|
||||
#
|
||||
TRANSFORMS:
|
||||
-
|
||||
|
||||
@@ -14,13 +14,14 @@ SOLVER:
|
||||
NUM_FOLDS: 1
|
||||
ACCU_STEP: 1
|
||||
EVAL_INTERVAL: 100
|
||||
RESCALE_LR: False
|
||||
#
|
||||
WORK_DIR: ./cache/save_data/sd15_512_sce_t2i_swift
|
||||
LOG_FILE: std_log.txt
|
||||
#
|
||||
FILE_SYSTEM:
|
||||
NAME: "ModelscopeFs"
|
||||
TEMP_DIR: "./cache/data"
|
||||
TEMP_DIR: "./cache/cache_data"
|
||||
#
|
||||
TUNER:
|
||||
-
|
||||
@@ -132,7 +133,7 @@ SOLVER:
|
||||
#
|
||||
OPTIMIZER:
|
||||
NAME: AdamW
|
||||
LEARNING_RATE: 0.064
|
||||
LEARNING_RATE: 0.0001
|
||||
BETAS: [ 0.9, 0.999 ]
|
||||
EPS: 1e-8
|
||||
WEIGHT_DECAY: 1e-2
|
||||
@@ -198,7 +199,7 @@ SOLVER:
|
||||
NUM_WORKERS: 4
|
||||
FILE_SYSTEM:
|
||||
NAME: "ModelscopeFs"
|
||||
TEMP_DIR: "./cache/data"
|
||||
TEMP_DIR: "./cache/cache_data"
|
||||
#
|
||||
TRANSFORMS:
|
||||
-
|
||||
|
||||
@@ -14,13 +14,14 @@ SOLVER:
|
||||
NUM_FOLDS: 1
|
||||
ACCU_STEP: 1
|
||||
EVAL_INTERVAL: 100
|
||||
RESCALE_LR: False
|
||||
#
|
||||
WORK_DIR: ./cache/save_data/sd15_512_textsce_t2i_swift
|
||||
LOG_FILE: std_log.txt
|
||||
#
|
||||
FILE_SYSTEM:
|
||||
NAME: "ModelscopeFs"
|
||||
TEMP_DIR: "./cache/data"
|
||||
TEMP_DIR: "./cache/cache_data"
|
||||
#
|
||||
TUNER:
|
||||
-
|
||||
@@ -140,7 +141,7 @@ SOLVER:
|
||||
#
|
||||
OPTIMIZER:
|
||||
NAME: AdamW
|
||||
LEARNING_RATE: 0.064
|
||||
LEARNING_RATE: 0.0001
|
||||
BETAS: [ 0.9, 0.999 ]
|
||||
EPS: 1e-8
|
||||
WEIGHT_DECAY: 1e-2
|
||||
@@ -206,7 +207,7 @@ SOLVER:
|
||||
NUM_WORKERS: 4
|
||||
FILE_SYSTEM:
|
||||
NAME: "ModelscopeFs"
|
||||
TEMP_DIR: "./cache/data"
|
||||
TEMP_DIR: "./cache/cache_data"
|
||||
#
|
||||
TRANSFORMS:
|
||||
-
|
||||
|
||||
@@ -14,13 +14,14 @@ SOLVER:
|
||||
NUM_FOLDS: 1
|
||||
ACCU_STEP: 1
|
||||
EVAL_INTERVAL: 100
|
||||
RESCALE_LR: False
|
||||
#
|
||||
WORK_DIR: ./cache/save_data/sd21_768_sce_t2i
|
||||
LOG_FILE: std_log.txt
|
||||
#
|
||||
FILE_SYSTEM:
|
||||
NAME: "ModelscopeFs"
|
||||
TEMP_DIR: "./cache/data"
|
||||
TEMP_DIR: "./cache/cache_data"
|
||||
#
|
||||
FREEZE:
|
||||
FREEZE_PART: [ "first_stage_model", "cond_stage_model", "model" ]
|
||||
@@ -130,7 +131,7 @@ SOLVER:
|
||||
#
|
||||
OPTIMIZER:
|
||||
NAME: AdamW
|
||||
LEARNING_RATE: 0.064
|
||||
LEARNING_RATE: 0.0001
|
||||
BETAS: [ 0.9, 0.999 ]
|
||||
EPS: 1e-8
|
||||
WEIGHT_DECAY: 1e-2
|
||||
@@ -196,7 +197,7 @@ SOLVER:
|
||||
NUM_WORKERS: 4
|
||||
FILE_SYSTEM:
|
||||
NAME: "ModelscopeFs"
|
||||
TEMP_DIR: "./cache/data"
|
||||
TEMP_DIR: "./cache/cache_data"
|
||||
#
|
||||
TRANSFORMS:
|
||||
-
|
||||
|
||||
@@ -14,13 +14,14 @@ SOLVER:
|
||||
NUM_FOLDS: 1
|
||||
ACCU_STEP: 1
|
||||
EVAL_INTERVAL: 100
|
||||
RESCALE_LR: False
|
||||
#
|
||||
WORK_DIR: ./cache/save_data/sd21_768_sce_t2i_swift
|
||||
LOG_FILE: std_log.txt
|
||||
#
|
||||
FILE_SYSTEM:
|
||||
NAME: "ModelscopeFs"
|
||||
TEMP_DIR: "./cache/data"
|
||||
TEMP_DIR: "./cache/cache_data"
|
||||
#
|
||||
TUNER:
|
||||
-
|
||||
@@ -128,7 +129,7 @@ SOLVER:
|
||||
#
|
||||
OPTIMIZER:
|
||||
NAME: AdamW
|
||||
LEARNING_RATE: 0.064
|
||||
LEARNING_RATE: 0.0001
|
||||
BETAS: [ 0.9, 0.999 ]
|
||||
EPS: 1e-8
|
||||
WEIGHT_DECAY: 1e-2
|
||||
@@ -194,7 +195,7 @@ SOLVER:
|
||||
NUM_WORKERS: 4
|
||||
FILE_SYSTEM:
|
||||
NAME: "ModelscopeFs"
|
||||
TEMP_DIR: "./cache/data"
|
||||
TEMP_DIR: "./cache/cache_data"
|
||||
#
|
||||
TRANSFORMS:
|
||||
-
|
||||
|
||||
@@ -14,13 +14,14 @@ SOLVER:
|
||||
NUM_FOLDS: 1
|
||||
ACCU_STEP: 1
|
||||
EVAL_INTERVAL: 100
|
||||
RESCALE_LR: False
|
||||
#
|
||||
WORK_DIR: ./cache/save_data/sdxl_1024_sce_t2i
|
||||
LOG_FILE: std_log.txt
|
||||
#
|
||||
FILE_SYSTEM:
|
||||
NAME: "ModelscopeFs"
|
||||
TEMP_DIR: "./cache/data"
|
||||
TEMP_DIR: "./cache/cache_data"
|
||||
#
|
||||
FREEZE:
|
||||
FREEZE_PART: [ "first_stage_model", "cond_stage_model", "model" ]
|
||||
@@ -247,7 +248,7 @@ SOLVER:
|
||||
#
|
||||
OPTIMIZER:
|
||||
NAME: AdamW
|
||||
LEARNING_RATE: 0.064
|
||||
LEARNING_RATE: 0.0001
|
||||
BETAS: [ 0.9, 0.999 ]
|
||||
EPS: 1e-8
|
||||
WEIGHT_DECAY: 1e-2
|
||||
@@ -315,7 +316,7 @@ SOLVER:
|
||||
NUM_WORKERS: 4
|
||||
FILE_SYSTEM:
|
||||
NAME: "ModelscopeFs"
|
||||
TEMP_DIR: "./cache/data"
|
||||
TEMP_DIR: "./cache/cache_data"
|
||||
#
|
||||
TRANSFORMS:
|
||||
- NAME: Select
|
||||
|
||||
@@ -14,13 +14,14 @@ SOLVER:
|
||||
NUM_FOLDS: 1
|
||||
ACCU_STEP: 1
|
||||
EVAL_INTERVAL: 100
|
||||
RESCALE_LR: False
|
||||
#
|
||||
WORK_DIR: ./cache/save_data/sdxl_1024_sce_t2i_datatxt
|
||||
LOG_FILE: std_log.txt
|
||||
#
|
||||
FILE_SYSTEM:
|
||||
NAME: "ModelscopeFs"
|
||||
TEMP_DIR: "./cache/data"
|
||||
TEMP_DIR: "./cache/cache_data"
|
||||
#
|
||||
FREEZE:
|
||||
FREEZE_PART: [ "first_stage_model", "cond_stage_model", "model" ]
|
||||
@@ -247,7 +248,7 @@ SOLVER:
|
||||
#
|
||||
OPTIMIZER:
|
||||
NAME: AdamW
|
||||
LEARNING_RATE: 0.064
|
||||
LEARNING_RATE: 0.0001
|
||||
BETAS: [ 0.9, 0.999 ]
|
||||
EPS: 1e-8
|
||||
WEIGHT_DECAY: 1e-2
|
||||
|
||||
@@ -14,13 +14,14 @@ SOLVER:
|
||||
NUM_FOLDS: 1
|
||||
ACCU_STEP: 1
|
||||
EVAL_INTERVAL: 100
|
||||
RESCALE_LR: False
|
||||
#
|
||||
WORK_DIR: ./cache/save_data/sdxl_1024_sce_t2i_swift
|
||||
LOG_FILE: std_log.txt
|
||||
#
|
||||
FILE_SYSTEM:
|
||||
NAME: "ModelscopeFs"
|
||||
TEMP_DIR: "./cache/data"
|
||||
TEMP_DIR: "./cache/cache_data"
|
||||
#
|
||||
TUNER:
|
||||
-
|
||||
@@ -245,7 +246,7 @@ SOLVER:
|
||||
#
|
||||
OPTIMIZER:
|
||||
NAME: AdamW
|
||||
LEARNING_RATE: 0.064
|
||||
LEARNING_RATE: 0.0001
|
||||
BETAS: [ 0.9, 0.999 ]
|
||||
EPS: 1e-8
|
||||
WEIGHT_DECAY: 1e-2
|
||||
@@ -313,7 +314,7 @@ SOLVER:
|
||||
NUM_WORKERS: 4
|
||||
FILE_SYSTEM:
|
||||
NAME: "ModelscopeFs"
|
||||
TEMP_DIR: "./cache/data"
|
||||
TEMP_DIR: "./cache/cache_data"
|
||||
#
|
||||
TRANSFORMS:
|
||||
- NAME: Select
|
||||
|
||||
@@ -14,13 +14,14 @@ SOLVER:
|
||||
NUM_FOLDS: 1
|
||||
ACCU_STEP: 1
|
||||
EVAL_INTERVAL: 100
|
||||
RESCALE_LR: False
|
||||
#
|
||||
WORK_DIR: ./cache/save_data/sdxl_1024_textsce_t2i_swift
|
||||
LOG_FILE: std_log.txt
|
||||
#
|
||||
FILE_SYSTEM:
|
||||
NAME: "ModelscopeFs"
|
||||
TEMP_DIR: "./cache/data"
|
||||
TEMP_DIR: "./cache/cache_data"
|
||||
#
|
||||
TUNER:
|
||||
-
|
||||
@@ -253,7 +254,7 @@ SOLVER:
|
||||
#
|
||||
OPTIMIZER:
|
||||
NAME: AdamW
|
||||
LEARNING_RATE: 0.064
|
||||
LEARNING_RATE: 0.0001
|
||||
BETAS: [ 0.9, 0.999 ]
|
||||
EPS: 1e-8
|
||||
WEIGHT_DECAY: 1e-2
|
||||
@@ -321,7 +322,7 @@ SOLVER:
|
||||
NUM_WORKERS: 4
|
||||
FILE_SYSTEM:
|
||||
NAME: "ModelscopeFs"
|
||||
TEMP_DIR: "./cache/data"
|
||||
TEMP_DIR: "./cache/cache_data"
|
||||
#
|
||||
TRANSFORMS:
|
||||
- NAME: Select
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,120 @@
|
||||
NAME: PIXART_ALPHA
|
||||
IS_DEFAULT: False
|
||||
DEFAULT_PARAS:
|
||||
PARAS:
|
||||
RESOLUTIONS: [[1024, 1024]]
|
||||
INPUT:
|
||||
IMAGE:
|
||||
ORIGINAL_SIZE_AS_TUPLE: [1024, 1024]
|
||||
TARGET_SIZE_AS_TUPLE: [1024, 1024]
|
||||
PROMPT: ""
|
||||
NEGATIVE_PROMPT: ""
|
||||
PROMPT_PREFIX: ""
|
||||
SAMPLE: ddim
|
||||
SAMPLE_STEPS: 20
|
||||
GUIDE_SCALE: 4.5
|
||||
GUIDE_RESCALE: 0.5
|
||||
DISCRETIZATION: trailing
|
||||
OUTPUT:
|
||||
LATENT:
|
||||
IMAGES:
|
||||
SEED:
|
||||
MODULES_PARAS:
|
||||
FIRST_STAGE_MODEL:
|
||||
FUNCTION:
|
||||
-
|
||||
NAME: encode
|
||||
DTYPE: float32
|
||||
INPUT: ["IMAGE"]
|
||||
-
|
||||
NAME: decode
|
||||
DTYPE: float32
|
||||
INPUT: ["LATENT"]
|
||||
PARAS:
|
||||
SCALE_FACTOR: 0.18215
|
||||
SIZE_FACTOR: 8
|
||||
DIFFUSION_MODEL:
|
||||
FUNCTION:
|
||||
-
|
||||
NAME: forward
|
||||
DTYPE: float16
|
||||
INPUT: ["SAMPLE_STEPS", "SAMPLE", "GUIDE_SCALE", "GUIDE_RESCALE", "DISCRETIZATION"]
|
||||
COND_STAGE_MODEL:
|
||||
FUNCTION:
|
||||
-
|
||||
NAME: encode
|
||||
DTYPE: float32
|
||||
INPUT: ["PROMPT"]
|
||||
#
|
||||
MODEL:
|
||||
PRETRAINED_MODEL:
|
||||
DECODER_BIAS: 0.5
|
||||
SCHEDULE:
|
||||
PARAMETERIZATION: "eps"
|
||||
TIMESTEPS: 1000
|
||||
ZERO_TERMINAL_SNR: False
|
||||
SCHEDULE_ARGS:
|
||||
"NAME": "linear"
|
||||
"BETA_MIN": 0.0001
|
||||
"BETA_MAX": 0.02
|
||||
#
|
||||
DIFFUSION_MODEL:
|
||||
NAME: PixArt
|
||||
PRETRAINED_MODEL: ms://AI-ModelScope/PixArt-alpha@PixArt-XL-2-1024-MS.pth
|
||||
INPUT_SIZE: 128
|
||||
PATCH_SIZE: 2
|
||||
IN_CHANNELS: 4
|
||||
HIDDEN_SIZE: 1152
|
||||
DEPTH: 28
|
||||
NUM_HEADS: 16
|
||||
MLP_RATIO: 4.0
|
||||
CLASS_DROPOUT_PROB: 0.1
|
||||
PRED_SIGMA: True
|
||||
DROP_PATH: 0.0
|
||||
WINDOW_DIZE: 0
|
||||
USE_REL_POS: False
|
||||
CAPTION_CHANNELS: 4096
|
||||
LEWEI_SCALE: 2
|
||||
MODEL_MAX_LENGTH: 120
|
||||
#
|
||||
FIRST_STAGE_MODEL:
|
||||
NAME: AutoencoderKL
|
||||
PRETRAINED_MODEL: ms://AI-ModelScope/stable-diffusion-2-base@512-base-ema.safetensors
|
||||
EMBED_DIM: 4
|
||||
IGNORE_KEYS: [ ]
|
||||
BATCH_SIZE: 1
|
||||
#
|
||||
ENCODER:
|
||||
NAME: Encoder
|
||||
CH: 128
|
||||
OUT_CH: 3
|
||||
NUM_RES_BLOCKS: 2
|
||||
IN_CHANNELS: 3
|
||||
ATTN_RESOLUTIONS: [ ]
|
||||
CH_MULT: [ 1, 2, 4, 4 ]
|
||||
Z_CHANNELS: 4
|
||||
DOUBLE_Z: True
|
||||
DROPOUT: 0.0
|
||||
RESAMP_WITH_CONV: True
|
||||
#
|
||||
DECODER:
|
||||
NAME: Decoder
|
||||
CH: 128
|
||||
OUT_CH: 3
|
||||
NUM_RES_BLOCKS: 2
|
||||
IN_CHANNELS: 3
|
||||
ATTN_RESOLUTIONS: [ ]
|
||||
CH_MULT: [ 1, 2, 4, 4 ]
|
||||
Z_CHANNELS: 4
|
||||
DROPOUT: 0.0
|
||||
RESAMP_WITH_CONV: True
|
||||
GIVE_PRE_END: False
|
||||
TANH_OUT: False
|
||||
#
|
||||
COND_STAGE_MODEL:
|
||||
NAME: T5EmbedderHF
|
||||
PRETRAINED_MODEL: ms://AI-ModelScope/PixArt-alpha@t5-v1_1-xxl/
|
||||
TOKENIZER_PATH: ms://AI-ModelScope/PixArt-alpha@t5-v1_1-xxl/
|
||||
LENGTH: 120
|
||||
CLEAN: heavy
|
||||
USE_GRAD: False
|
||||
@@ -0,0 +1,148 @@
|
||||
NAME: SD3
|
||||
IS_DEFAULT: False
|
||||
DEFAULT_PARAS:
|
||||
PARAS:
|
||||
RESOLUTIONS: [[1024, 1024]]
|
||||
INPUT:
|
||||
IMAGE:
|
||||
ORIGINAL_SIZE_AS_TUPLE: [1024, 1024]
|
||||
TARGET_SIZE_AS_TUPLE: [1024, 1024]
|
||||
PROMPT: ""
|
||||
NEGATIVE_PROMPT: ""
|
||||
PROMPT_PREFIX: ""
|
||||
SAMPLE: euler
|
||||
SAMPLE_STEPS: 28
|
||||
GUIDE_SCALE: 5.0
|
||||
GUIDE_RESCALE: 0.0
|
||||
DISCRETIZATION: trailing
|
||||
OUTPUT:
|
||||
LATENT:
|
||||
IMAGES:
|
||||
SEED:
|
||||
MODULES_PARAS:
|
||||
FIRST_STAGE_MODEL:
|
||||
FUNCTION:
|
||||
-
|
||||
NAME: encode
|
||||
DTYPE: float32
|
||||
INPUT: ["IMAGE"]
|
||||
-
|
||||
NAME: decode
|
||||
DTYPE: float32
|
||||
INPUT: ["LATENT"]
|
||||
PARAS:
|
||||
SCALE_FACTOR: 1.5305
|
||||
SHIFT_FACTOR: 0.0609
|
||||
SIZE_FACTOR: 8
|
||||
DIFFUSION_MODEL:
|
||||
FUNCTION:
|
||||
-
|
||||
NAME: forward
|
||||
DTYPE: float16
|
||||
INPUT: ["SAMPLE_STEPS", "SAMPLE", "GUIDE_SCALE", "GUIDE_RESCALE", "DISCRETIZATION"]
|
||||
COND_STAGE_MODEL:
|
||||
FUNCTION:
|
||||
-
|
||||
NAME: encode
|
||||
DTYPE: float32
|
||||
INPUT: ["PROMPT"]
|
||||
#
|
||||
MODEL:
|
||||
PRETRAINED_MODEL:
|
||||
SCHEDULE:
|
||||
PARAMETERIZATION: rf
|
||||
TIMESTEPS: 1000
|
||||
MIN_SNR_GAMMA:
|
||||
ZERO_TERMINAL_SNR: False
|
||||
PRETRAINED_MODEL:
|
||||
IGNORE_KEYS: [ ]
|
||||
SCALE_FACTOR: 1.5305
|
||||
SHIFT_FACTOR: 0.0609
|
||||
DEFAULT_N_PROMPT:
|
||||
SCHEDULE_ARGS:
|
||||
"NAME": "shifted"
|
||||
"SHIFT": 3
|
||||
T_WEIGHT: uniform
|
||||
#
|
||||
DIFFUSION_MODEL:
|
||||
NAME: MMDiT
|
||||
PRETRAINED_MODEL: ms://AI-ModelScope/stable-diffusion-3-medium@sd3_medium.safetensors
|
||||
IGNORE_KEYS: '^first_stage_model.'
|
||||
IN_CHANNELS: 16
|
||||
PATCH_SIZE: 2
|
||||
OUT_CHANNELS: 16
|
||||
DEPTH: 24
|
||||
INPUT_SIZE:
|
||||
ADM_IN_CHANNELS: 2048
|
||||
CONTEXT_EMBEDDER_CONFIG: { 'target': 'torch.nn.Linear', 'params': { 'in_features': 4096, 'out_features': 1536 } }
|
||||
NUM_PATCHES: 36864
|
||||
POS_EMBED_MAX_SIZE: 192
|
||||
POS_EMBED_SCALING_FACTOR:
|
||||
USE_CHECKPOINT: True
|
||||
#
|
||||
FIRST_STAGE_MODEL:
|
||||
NAME: AutoencoderKL
|
||||
PRETRAINED_MODEL: ms://AI-ModelScope/stable-diffusion-3-medium@sd3_medium.safetensors
|
||||
EMBED_DIM: 16
|
||||
IGNORE_KEYS: '^model.diffusion_model.'
|
||||
BATCH_SIZE: 1
|
||||
USE_CONV: False
|
||||
#
|
||||
ENCODER:
|
||||
NAME: Encoder
|
||||
CH: 128
|
||||
OUT_CH: 3
|
||||
NUM_RES_BLOCKS: 2
|
||||
IN_CHANNELS: 3
|
||||
ATTN_RESOLUTIONS: [ ]
|
||||
CH_MULT: [ 1, 2, 4, 4 ]
|
||||
Z_CHANNELS: 16
|
||||
DOUBLE_Z: True
|
||||
DROPOUT: 0.0
|
||||
RESAMP_WITH_CONV: True
|
||||
#
|
||||
DECODER:
|
||||
NAME: Decoder
|
||||
CH: 128
|
||||
OUT_CH: 3
|
||||
NUM_RES_BLOCKS: 2
|
||||
IN_CHANNELS: 3
|
||||
ATTN_RESOLUTIONS: [ ]
|
||||
CH_MULT: [ 1, 2, 4, 4 ]
|
||||
Z_CHANNELS: 16
|
||||
DROPOUT: 0.0
|
||||
RESAMP_WITH_CONV: True
|
||||
GIVE_PRE_END: False
|
||||
TANH_OUT: False
|
||||
#
|
||||
COND_STAGE_MODEL:
|
||||
NAME: SD3TextEmbedder
|
||||
P_ZERO: 0.0
|
||||
CLIP_L:
|
||||
NAME: FrozenCLIPEmbedder2
|
||||
PRETRAINED_MODEL: ms://AI-ModelScope/stable-diffusion-3-medium-diffusers@text_encoder
|
||||
TOKENIZER_PATH: ms://AI-ModelScope/stable-diffusion-3-medium-diffusers@tokenizer
|
||||
MAX_LENGTH: 77
|
||||
FREEZE: True
|
||||
LAYER: penultimate
|
||||
RETURN_POOLED: True
|
||||
USE_FINAL_LAYER_NORM: False
|
||||
IS_TRAINABLE: False
|
||||
CLIP_G:
|
||||
NAME: FrozenCLIPEmbedder2
|
||||
PRETRAINED_MODEL: ms://AI-ModelScope/stable-diffusion-3-medium-diffusers@text_encoder_2
|
||||
TOKENIZER_PATH: ms://AI-ModelScope/stable-diffusion-3-medium-diffusers@tokenizer_2
|
||||
MAX_LENGTH: 77
|
||||
FREEZE: True
|
||||
LAYER: penultimate
|
||||
RETURN_POOLED: True
|
||||
USE_FINAL_LAYER_NORM: False
|
||||
IS_TRAINABLE: False
|
||||
T5_XXL:
|
||||
NAME: T5EmbedderHF
|
||||
PRETRAINED_MODEL: ms://AI-ModelScope/stable-diffusion-3-medium-diffusers@text_encoder_3
|
||||
TOKENIZER_PATH: ms://AI-ModelScope/stable-diffusion-3-medium-diffusers@tokenizer_3
|
||||
LENGTH: 256
|
||||
CLEAN: whitespace
|
||||
USE_GRAD: False
|
||||
T5_DTYPE: float16
|
||||
@@ -0,0 +1,264 @@
|
||||
ENV:
|
||||
BACKEND: nccl
|
||||
META:
|
||||
VERSION: 'PIXART_ALPHA'
|
||||
DESCRIPTION: "PIXART ALPHA"
|
||||
IS_DEFAULT: False
|
||||
IS_SHARE: True
|
||||
INFERENCE_PARAS:
|
||||
INFERENCE_BATCH_SIZE: 1
|
||||
INFERENCE_PREFIX: ""
|
||||
DEFAULT_SAMPLER: "ddim"
|
||||
DEFAULT_SAMPLE_STEPS: 20
|
||||
INFERENCE_N_PROMPT: ""
|
||||
RESOLUTION: [1024, 1024]
|
||||
PARAS:
|
||||
-
|
||||
TRAIN_BATCH_SIZE: 2
|
||||
TRAIN_PREFIX: ""
|
||||
TRAIN_N_PROMPT: ""
|
||||
RESOLUTION: [1024, 1024]
|
||||
MEMORY: 29000
|
||||
EPOCHS: 50
|
||||
SAVE_INTERVAL: 25
|
||||
EPSEC: 0.818
|
||||
LEARNING_RATE: 0.0001
|
||||
IS_DEFAULT: False
|
||||
TUNER: FULL
|
||||
-
|
||||
TRAIN_BATCH_SIZE: 2
|
||||
TRAIN_PREFIX: ""
|
||||
TRAIN_N_PROMPT: ""
|
||||
RESOLUTION: [1024, 1024]
|
||||
MEMORY: 29000
|
||||
EPOCHS: 50
|
||||
SAVE_INTERVAL: 25
|
||||
EPSEC: 0.818
|
||||
LEARNING_RATE: 0.0001
|
||||
IS_DEFAULT: False
|
||||
TUNER: LORA
|
||||
#
|
||||
TUNERS:
|
||||
LORA:
|
||||
-
|
||||
NAME: SwiftLoRA
|
||||
R: 256
|
||||
LORA_ALPHA: 256
|
||||
LORA_DROPOUT: 0.0
|
||||
BIAS: "none"
|
||||
TARGET_MODULES: "model.*(.q|.k|.v|.o|mlp.fc1|mlp.fc2)$"
|
||||
#
|
||||
SOLVER:
|
||||
NAME: LatentDiffusionSolver
|
||||
RESUME_FROM:
|
||||
LOAD_MODEL_ONLY: True
|
||||
USE_FSDP: False
|
||||
SHARDING_STRATEGY:
|
||||
USE_AMP: True
|
||||
DTYPE: float16
|
||||
CHANNELS_LAST: True
|
||||
MAX_STEPS: 1000
|
||||
MAX_EPOCHS: -1
|
||||
NUM_FOLDS: 1
|
||||
ACCU_STEP: 1
|
||||
EVAL_INTERVAL: -1
|
||||
RESCALE_LR: False
|
||||
#
|
||||
WORK_DIR: cache/scepter_ui/self_train/dit/pixart_alpha_pro
|
||||
LOG_FILE: std_log.txt
|
||||
#
|
||||
FILE_SYSTEM:
|
||||
NAME: "ModelscopeFs"
|
||||
TEMP_DIR: "./cache/cache_data"
|
||||
#
|
||||
FREEZE:
|
||||
#
|
||||
TUNER:
|
||||
#
|
||||
MODEL:
|
||||
NAME: LatentDiffusionPixart
|
||||
PARAMETERIZATION: eps
|
||||
TIMESTEPS: 1000
|
||||
MIN_SNR_GAMMA:
|
||||
ZERO_TERMINAL_SNR: False
|
||||
PRETRAINED_MODEL:
|
||||
IGNORE_KEYS: [ ]
|
||||
SCALE_FACTOR: 0.18215
|
||||
SIZE_FACTOR: 8
|
||||
DECODER_BIAS: 0.5
|
||||
DEFAULT_N_PROMPT:
|
||||
SCHEDULE_ARGS:
|
||||
"NAME": "linear"
|
||||
"BETA_MIN": 0.0001
|
||||
"BETA_MAX": 0.02
|
||||
USE_EMA: False
|
||||
LOAD_REFINER: False
|
||||
#
|
||||
DIFFUSION_MODEL:
|
||||
NAME: PixArt
|
||||
PRETRAINED_MODEL: ms://AI-ModelScope/PixArt-alpha@PixArt-XL-2-1024-MS.pth
|
||||
INPUT_SIZE: 128
|
||||
PATCH_SIZE: 2
|
||||
IN_CHANNELS: 4
|
||||
HIDDEN_SIZE: 1152
|
||||
DEPTH: 28
|
||||
NUM_HEADS: 16
|
||||
MLP_RATIO: 4.0
|
||||
CLASS_DROPOUT_PROB: 0.1
|
||||
PRED_SIGMA: True
|
||||
DROP_PATH: 0.0
|
||||
WINDOW_DIZE: 0
|
||||
USE_REL_POS: False
|
||||
CAPTION_CHANNELS: 4096
|
||||
LEWEI_SCALE: 2
|
||||
MODEL_MAX_LENGTH: 120
|
||||
#
|
||||
FIRST_STAGE_MODEL:
|
||||
NAME: AutoencoderKL
|
||||
PRETRAINED_MODEL: ms://AI-ModelScope/stable-diffusion-2-base@512-base-ema.safetensors
|
||||
EMBED_DIM: 4
|
||||
IGNORE_KEYS: [ ]
|
||||
BATCH_SIZE: 1
|
||||
#
|
||||
ENCODER:
|
||||
NAME: Encoder
|
||||
CH: 128
|
||||
OUT_CH: 3
|
||||
NUM_RES_BLOCKS: 2
|
||||
IN_CHANNELS: 3
|
||||
ATTN_RESOLUTIONS: [ ]
|
||||
CH_MULT: [ 1, 2, 4, 4 ]
|
||||
Z_CHANNELS: 4
|
||||
DOUBLE_Z: True
|
||||
DROPOUT: 0.0
|
||||
RESAMP_WITH_CONV: True
|
||||
#
|
||||
DECODER:
|
||||
NAME: Decoder
|
||||
CH: 128
|
||||
OUT_CH: 3
|
||||
NUM_RES_BLOCKS: 2
|
||||
IN_CHANNELS: 3
|
||||
ATTN_RESOLUTIONS: [ ]
|
||||
CH_MULT: [ 1, 2, 4, 4 ]
|
||||
Z_CHANNELS: 4
|
||||
DROPOUT: 0.0
|
||||
RESAMP_WITH_CONV: True
|
||||
GIVE_PRE_END: False
|
||||
TANH_OUT: False
|
||||
#
|
||||
COND_STAGE_MODEL:
|
||||
NAME: T5EmbedderHF
|
||||
PRETRAINED_MODEL: ms://AI-ModelScope/PixArt-alpha@t5-v1_1-xxl/
|
||||
TOKENIZER_PATH: ms://AI-ModelScope/PixArt-alpha@t5-v1_1-xxl/
|
||||
LENGTH: 120
|
||||
CLEAN: heavy
|
||||
USE_GRAD: False
|
||||
#
|
||||
LOSS:
|
||||
NAME: ReconstructLoss
|
||||
LOSS_TYPE: l2
|
||||
#
|
||||
SAMPLE_ARGS:
|
||||
SAMPLER: ddim
|
||||
SAMPLE_STEPS: 20
|
||||
SEED: 2024
|
||||
GUIDE_SCALE: 4.5
|
||||
GUIDE_RESCALE: 0.5
|
||||
DISCRETIZATION: trailing
|
||||
RUN_TRAIN_N: False
|
||||
#
|
||||
OPTIMIZER:
|
||||
NAME: AdamW
|
||||
LEARNING_RATE: 0.0001
|
||||
BETAS: [ 0.9, 0.999 ]
|
||||
EPS: 1e-8
|
||||
WEIGHT_DECAY: 1e-2
|
||||
AMSGRAD: False
|
||||
#
|
||||
TRAIN_DATA:
|
||||
NAME: ImageTextPairMSDataset
|
||||
MODE: train
|
||||
MS_DATASET_NAME: style_custom_dataset
|
||||
MS_DATASET_NAMESPACE: damo
|
||||
MS_DATASET_SUBNAME: 3D
|
||||
PROMPT_PREFIX: ""
|
||||
MS_DATASET_SPLIT: train
|
||||
MS_REMAP_KEYS: { 'Image:FILE': 'Target:FILE' }
|
||||
REPLACE_STYLE: False
|
||||
PIN_MEMORY: True
|
||||
BATCH_SIZE: 1
|
||||
NUM_WORKERS: 4
|
||||
SAMPLER:
|
||||
NAME: LoopSampler
|
||||
TRANSFORMS:
|
||||
- NAME: LoadImageFromFile
|
||||
RGB_ORDER: RGB
|
||||
BACKEND: pillow
|
||||
- NAME: FlexibleResize
|
||||
INTERPOLATION: bilinear
|
||||
SIZE: [ 1024, 1024 ]
|
||||
INPUT_KEY: [ 'img' ]
|
||||
OUTPUT_KEY: [ 'img' ]
|
||||
BACKEND: pillow
|
||||
- NAME: FlexibleCenterCrop
|
||||
SIZE: [ 1024, 1024 ]
|
||||
INPUT_KEY: [ 'img' ]
|
||||
OUTPUT_KEY: [ 'img' ]
|
||||
BACKEND: pillow
|
||||
- NAME: ImageToTensor
|
||||
INPUT_KEY: [ 'img' ]
|
||||
OUTPUT_KEY: [ 'img' ]
|
||||
BACKEND: pillow
|
||||
- NAME: Normalize
|
||||
MEAN: [ 0.5, 0.5, 0.5 ]
|
||||
STD: [ 0.5, 0.5, 0.5 ]
|
||||
INPUT_KEY: [ 'img' ]
|
||||
OUTPUT_KEY: [ 'image' ]
|
||||
BACKEND: torchvision
|
||||
- NAME: Select
|
||||
KEYS: [ 'image', 'prompt' ]
|
||||
META_KEYS: [ 'data_key' ]
|
||||
#
|
||||
EVAL_DATA:
|
||||
NAME: Text2ImageDataset
|
||||
MODE: eval
|
||||
PROMPT_FILE:
|
||||
PROMPT_DATA: [ "a boy wearing a jacket", "a dog running on the lawn" ]
|
||||
IMAGE_SIZE: [ 1024, 1024 ]
|
||||
FIELDS: [ "prompt" ]
|
||||
DELIMITER: '#;#'
|
||||
PROMPT_PREFIX: ''
|
||||
PIN_MEMORY: True
|
||||
BATCH_SIZE: 1
|
||||
NUM_WORKERS: 4
|
||||
TRANSFORMS:
|
||||
- NAME: Select
|
||||
KEYS: [ 'index', 'prompt' ]
|
||||
META_KEYS: [ 'image_size' ]
|
||||
#
|
||||
TRAIN_HOOKS:
|
||||
-
|
||||
NAME: BackwardHook
|
||||
PRIORITY: 0
|
||||
-
|
||||
NAME: LogHook
|
||||
LOG_INTERVAL: 10
|
||||
SHOW_GPU_MEM: True
|
||||
-
|
||||
NAME: TensorboardLogHook
|
||||
-
|
||||
NAME: CheckpointHook
|
||||
INTERVAL: 10000
|
||||
PRIORITY: 200
|
||||
SAVE_LAST: True
|
||||
SAVE_NAME_PREFIX: 'step'
|
||||
DISABLE_SNAPSHOT: True
|
||||
#
|
||||
EVAL_HOOKS:
|
||||
-
|
||||
NAME: ProbeDataHook
|
||||
PROB_INTERVAL: 100
|
||||
SAVE_LAST: True
|
||||
SAVE_NAME_PREFIX: 'step'
|
||||
SAVE_PROBE_PREFIX: 'image'
|
||||
@@ -0,0 +1,284 @@
|
||||
ENV:
|
||||
BACKEND: nccl
|
||||
META:
|
||||
VERSION: 'SD3'
|
||||
DESCRIPTION: "SD3 ALPHA"
|
||||
IS_DEFAULT: False
|
||||
IS_SHARE: True
|
||||
INFERENCE_PARAS:
|
||||
INFERENCE_BATCH_SIZE: 1
|
||||
INFERENCE_PREFIX: ""
|
||||
DEFAULT_SAMPLER: "euler"
|
||||
DEFAULT_SAMPLE_STEPS: 28
|
||||
INFERENCE_N_PROMPT: ""
|
||||
RESOLUTION: [1024, 1024]
|
||||
PARAS:
|
||||
-
|
||||
TRAIN_BATCH_SIZE: 2
|
||||
TRAIN_PREFIX: ""
|
||||
TRAIN_N_PROMPT: ""
|
||||
RESOLUTION: [1024, 1024]
|
||||
MEMORY: 29000
|
||||
EPOCHS: 50
|
||||
SAVE_INTERVAL: 25
|
||||
EPSEC: 0.818
|
||||
LEARNING_RATE: 1e-5
|
||||
IS_DEFAULT: False
|
||||
TUNER: FULL
|
||||
-
|
||||
TRAIN_BATCH_SIZE: 2
|
||||
TRAIN_PREFIX: ""
|
||||
TRAIN_N_PROMPT: ""
|
||||
RESOLUTION: [1024, 1024]
|
||||
MEMORY: 29000
|
||||
EPOCHS: 50
|
||||
SAVE_INTERVAL: 25
|
||||
EPSEC: 0.818
|
||||
LEARNING_RATE: 5e-5
|
||||
IS_DEFAULT: True
|
||||
TUNER: LORA
|
||||
#
|
||||
TUNERS:
|
||||
LORA:
|
||||
-
|
||||
NAME: SwiftLoRA
|
||||
R: 128
|
||||
LORA_ALPHA: 128
|
||||
LORA_DROPOUT: 0.0
|
||||
BIAS: "none"
|
||||
TARGET_MODULES: "model.*(.attn.qkv|.attn.proj|mlp.fc1|mlp.fc2)$"
|
||||
#
|
||||
SOLVER:
|
||||
NAME: LatentDiffusionSolver
|
||||
RESUME_FROM:
|
||||
LOAD_MODEL_ONLY: True
|
||||
USE_FSDP: False
|
||||
SHARDING_STRATEGY:
|
||||
USE_AMP: True
|
||||
DTYPE: float16
|
||||
CHANNELS_LAST: True
|
||||
MAX_STEPS: 1000
|
||||
MAX_EPOCHS: -1
|
||||
NUM_FOLDS: 1
|
||||
ACCU_STEP: 1
|
||||
EVAL_INTERVAL: -1
|
||||
RESCALE_LR: False
|
||||
#
|
||||
WORK_DIR: cache/scepter_ui/self_train/dit/sd3
|
||||
LOG_FILE: std_log.txt
|
||||
#
|
||||
FILE_SYSTEM:
|
||||
- NAME: "ModelscopeFs"
|
||||
TEMP_DIR: "./cache/cache_data"
|
||||
#
|
||||
FREEZE:
|
||||
#
|
||||
TUNER:
|
||||
#
|
||||
MODEL:
|
||||
NAME: LatentDiffusionSD3
|
||||
PARAMETERIZATION: rf
|
||||
TIMESTEPS: 1000
|
||||
MIN_SNR_GAMMA:
|
||||
ZERO_TERMINAL_SNR: False
|
||||
PRETRAINED_MODEL:
|
||||
IGNORE_KEYS: [ ]
|
||||
SCALE_FACTOR: 1.5305
|
||||
SHIFT_FACTOR: 0.0609
|
||||
DEFAULT_N_PROMPT:
|
||||
SCHEDULE_ARGS:
|
||||
"NAME": "shifted"
|
||||
"SHIFT": 3
|
||||
USE_EMA: False
|
||||
T_WEIGHT: uniform
|
||||
#
|
||||
DIFFUSION_MODEL:
|
||||
NAME: MMDiT
|
||||
PRETRAINED_MODEL: ms://AI-ModelScope/stable-diffusion-3-medium@sd3_medium.safetensors
|
||||
IGNORE_KEYS: '^first_stage_model.'
|
||||
IN_CHANNELS: 16
|
||||
PATCH_SIZE: 2
|
||||
OUT_CHANNELS: 16
|
||||
DEPTH: 24
|
||||
INPUT_SIZE:
|
||||
ADM_IN_CHANNELS: 2048
|
||||
CONTEXT_EMBEDDER_CONFIG: { 'target': 'torch.nn.Linear', 'params': { 'in_features': 4096, 'out_features': 1536 } }
|
||||
NUM_PATCHES: 36864
|
||||
POS_EMBED_MAX_SIZE: 192
|
||||
POS_EMBED_SCALING_FACTOR:
|
||||
USE_CHECKPOINT: True
|
||||
#
|
||||
FIRST_STAGE_MODEL:
|
||||
NAME: AutoencoderKL
|
||||
PRETRAINED_MODEL: ms://AI-ModelScope/stable-diffusion-3-medium@sd3_medium.safetensors
|
||||
EMBED_DIM: 16
|
||||
IGNORE_KEYS: '^model.diffusion_model.'
|
||||
BATCH_SIZE: 1
|
||||
USE_CONV: False
|
||||
#
|
||||
ENCODER:
|
||||
NAME: Encoder
|
||||
CH: 128
|
||||
OUT_CH: 3
|
||||
NUM_RES_BLOCKS: 2
|
||||
IN_CHANNELS: 3
|
||||
ATTN_RESOLUTIONS: [ ]
|
||||
CH_MULT: [ 1, 2, 4, 4 ]
|
||||
Z_CHANNELS: 16
|
||||
DOUBLE_Z: True
|
||||
DROPOUT: 0.0
|
||||
RESAMP_WITH_CONV: True
|
||||
#
|
||||
DECODER:
|
||||
NAME: Decoder
|
||||
CH: 128
|
||||
OUT_CH: 3
|
||||
NUM_RES_BLOCKS: 2
|
||||
IN_CHANNELS: 3
|
||||
ATTN_RESOLUTIONS: [ ]
|
||||
CH_MULT: [ 1, 2, 4, 4 ]
|
||||
Z_CHANNELS: 16
|
||||
DROPOUT: 0.0
|
||||
RESAMP_WITH_CONV: True
|
||||
GIVE_PRE_END: False
|
||||
TANH_OUT: False
|
||||
#
|
||||
COND_STAGE_MODEL:
|
||||
NAME: SD3TextEmbedder
|
||||
P_ZERO: 0.0
|
||||
CLIP_L:
|
||||
NAME: FrozenCLIPEmbedder2
|
||||
PRETRAINED_MODEL: ms://AI-ModelScope/stable-diffusion-3-medium-diffusers@text_encoder
|
||||
TOKENIZER_PATH: ms://AI-ModelScope/stable-diffusion-3-medium-diffusers@tokenizer
|
||||
MAX_LENGTH: 77
|
||||
FREEZE: True
|
||||
LAYER: penultimate
|
||||
RETURN_POOLED: True
|
||||
USE_FINAL_LAYER_NORM: False
|
||||
IS_TRAINABLE: False
|
||||
CLIP_G:
|
||||
NAME: FrozenCLIPEmbedder2
|
||||
PRETRAINED_MODEL: ms://AI-ModelScope/stable-diffusion-3-medium-diffusers@text_encoder_2
|
||||
TOKENIZER_PATH: ms://AI-ModelScope/stable-diffusion-3-medium-diffusers@tokenizer_2
|
||||
MAX_LENGTH: 77
|
||||
FREEZE: True
|
||||
LAYER: penultimate
|
||||
RETURN_POOLED: True
|
||||
USE_FINAL_LAYER_NORM: False
|
||||
IS_TRAINABLE: False
|
||||
T5_XXL:
|
||||
NAME: T5EmbedderHF
|
||||
PRETRAINED_MODEL: ms://AI-ModelScope/stable-diffusion-3-medium-diffusers@text_encoder_3
|
||||
TOKENIZER_PATH: ms://AI-ModelScope/stable-diffusion-3-medium-diffusers@tokenizer_3
|
||||
LENGTH: 256
|
||||
CLEAN: whitespace
|
||||
USE_GRAD: False
|
||||
T5_DTYPE: float16
|
||||
#
|
||||
LOSS:
|
||||
NAME: ReconstructLoss
|
||||
LOSS_TYPE: l2
|
||||
#
|
||||
SAMPLE_ARGS:
|
||||
SAMPLER: euler
|
||||
SAMPLE_STEPS: 28
|
||||
SEED: 1749023094
|
||||
GUIDE_SCALE: 5.0
|
||||
GUIDE_RESCALE: 0.0
|
||||
DISCRETIZATION: trailing
|
||||
RUN_TRAIN_N: False
|
||||
#
|
||||
OPTIMIZER:
|
||||
NAME: AdamW
|
||||
LEARNING_RATE: 5e-5
|
||||
BETAS: [ 0.9, 0.999 ]
|
||||
EPS: 1e-8
|
||||
WEIGHT_DECAY: 1e-2
|
||||
AMSGRAD: False
|
||||
#
|
||||
TRAIN_DATA:
|
||||
NAME: ImageTextPairMSDataset
|
||||
MODE: train
|
||||
MS_DATASET_NAME: style_custom_dataset
|
||||
MS_DATASET_NAMESPACE: damo
|
||||
MS_DATASET_SUBNAME: 3D
|
||||
PROMPT_PREFIX: ""
|
||||
MS_DATASET_SPLIT: train
|
||||
MS_REMAP_KEYS: { 'Image:FILE': 'Target:FILE' }
|
||||
REPLACE_STYLE: False
|
||||
PIN_MEMORY: True
|
||||
BATCH_SIZE: 1
|
||||
NUM_WORKERS: 4
|
||||
SAMPLER:
|
||||
NAME: LoopSampler
|
||||
TRANSFORMS:
|
||||
- NAME: LoadImageFromFile
|
||||
RGB_ORDER: RGB
|
||||
BACKEND: pillow
|
||||
- NAME: FlexibleResize
|
||||
INTERPOLATION: bilinear
|
||||
SIZE: [ 1024, 1024 ]
|
||||
INPUT_KEY: [ 'img' ]
|
||||
OUTPUT_KEY: [ 'img' ]
|
||||
BACKEND: pillow
|
||||
- NAME: FlexibleCenterCrop
|
||||
SIZE: [ 1024, 1024 ]
|
||||
INPUT_KEY: [ 'img' ]
|
||||
OUTPUT_KEY: [ 'img' ]
|
||||
BACKEND: pillow
|
||||
- NAME: ImageToTensor
|
||||
INPUT_KEY: [ 'img' ]
|
||||
OUTPUT_KEY: [ 'img' ]
|
||||
BACKEND: pillow
|
||||
- NAME: Normalize
|
||||
MEAN: [ 0.5, 0.5, 0.5 ]
|
||||
STD: [ 0.5, 0.5, 0.5 ]
|
||||
INPUT_KEY: [ 'img' ]
|
||||
OUTPUT_KEY: [ 'image' ]
|
||||
BACKEND: torchvision
|
||||
- NAME: Select
|
||||
KEYS: [ 'image', 'prompt' ]
|
||||
META_KEYS: [ 'data_key' ]
|
||||
#
|
||||
EVAL_DATA:
|
||||
NAME: Text2ImageDataset
|
||||
MODE: eval
|
||||
PROMPT_FILE:
|
||||
PROMPT_DATA: [ "a cat holds a blackboard that writes \"hello world\"", "a dog running on the lawn" ]
|
||||
IMAGE_SIZE: [ 1024, 1024 ]
|
||||
FIELDS: [ "prompt" ]
|
||||
DELIMITER: '#;#'
|
||||
PROMPT_PREFIX: ''
|
||||
PIN_MEMORY: True
|
||||
BATCH_SIZE: 1
|
||||
NUM_WORKERS: 4
|
||||
TRANSFORMS:
|
||||
- NAME: Select
|
||||
KEYS: [ 'index', 'prompt' ]
|
||||
META_KEYS: [ 'image_size' ]
|
||||
#
|
||||
TRAIN_HOOKS:
|
||||
-
|
||||
NAME: BackwardHook
|
||||
PRIORITY: 1
|
||||
-
|
||||
NAME: LogHook
|
||||
LOG_INTERVAL: 10
|
||||
SHOW_GPU_MEM: True
|
||||
-
|
||||
NAME: TensorboardLogHook
|
||||
-
|
||||
NAME: CheckpointHook
|
||||
INTERVAL: 10000
|
||||
PRIORITY: 200
|
||||
SAVE_LAST: True
|
||||
SAVE_NAME_PREFIX: 'step'
|
||||
DISABLE_SNAPSHOT: True
|
||||
#
|
||||
EVAL_HOOKS:
|
||||
-
|
||||
NAME: ProbeDataHook
|
||||
PROB_INTERVAL: 100
|
||||
SAVE_LAST: True
|
||||
SAVE_NAME_PREFIX: 'step'
|
||||
SAVE_PROBE_PREFIX: 'image'
|
||||
@@ -142,13 +142,14 @@ SOLVER:
|
||||
NUM_FOLDS: 1
|
||||
ACCU_STEP: 1
|
||||
EVAL_INTERVAL: -1
|
||||
RESCALE_LR: False
|
||||
#
|
||||
WORK_DIR:
|
||||
LOG_FILE: std_log.txt
|
||||
#
|
||||
FILE_SYSTEM:
|
||||
NAME: "ModelscopeFs"
|
||||
TEMP_DIR: "./cache/data"
|
||||
TEMP_DIR: "./cache/cache_data"
|
||||
#
|
||||
FREEZE:
|
||||
#
|
||||
@@ -258,7 +259,7 @@ SOLVER:
|
||||
#
|
||||
OPTIMIZER:
|
||||
NAME: AdamW
|
||||
LEARNING_RATE: 0.064
|
||||
LEARNING_RATE: 0.0001
|
||||
BETAS: [ 0.9, 0.999 ]
|
||||
EPS: 1e-8
|
||||
WEIGHT_DECAY: 1e-2
|
||||
@@ -267,7 +268,7 @@ SOLVER:
|
||||
TRAIN_DATA:
|
||||
NAME: ImageTextPairMSDataset
|
||||
MODE: train
|
||||
MS_DATASET_NAME: cache/save_data/delogo/
|
||||
MS_DATASET_NAME: ./cache/cache_data/delogo/
|
||||
MS_DATASET_NAMESPACE: ""
|
||||
MS_DATASET_SPLIT: "train"
|
||||
MS_DATASET_SUBNAME: ""
|
||||
|
||||
@@ -147,6 +147,7 @@ SOLVER:
|
||||
MAX_EPOCHS: -1
|
||||
# NUM_FOLDS DESCRIPTION: Num folds for training. TYPE: int default: 1
|
||||
NUM_FOLDS: 1
|
||||
RESCALE_LR: False
|
||||
#
|
||||
EVAL_INTERVAL: -1
|
||||
# WORK_DIR DESCRIPTION: Save dir of the training log or model. TYPE: str default: ''
|
||||
@@ -155,7 +156,7 @@ SOLVER:
|
||||
LOG_FILE: std_log.txt
|
||||
FILE_SYSTEM:
|
||||
NAME: "ModelscopeFs"
|
||||
TEMP_DIR: "./cache/data"
|
||||
TEMP_DIR: "./cache/cache_data"
|
||||
TUNER:
|
||||
# MODEL DESCRIPTION: TYPE: default: ''
|
||||
MODEL:
|
||||
@@ -538,7 +539,7 @@ SOLVER:
|
||||
OPTIMIZER:
|
||||
# NAME DESCRIPTION: TYPE: default: ''
|
||||
NAME: AdamW
|
||||
LEARNING_RATE: 0.0064
|
||||
LEARNING_RATE: 0.00001
|
||||
EPS: 1e-8
|
||||
AMSGRAD: False
|
||||
#
|
||||
|
||||
@@ -137,13 +137,14 @@ SOLVER:
|
||||
NUM_FOLDS: 1
|
||||
ACCU_STEP: 1
|
||||
EVAL_INTERVAL: -1
|
||||
RESCALE_LR: False
|
||||
#
|
||||
WORK_DIR:
|
||||
LOG_FILE: std_log.txt
|
||||
#
|
||||
FILE_SYSTEM:
|
||||
NAME: "ModelscopeFs"
|
||||
TEMP_DIR: "./cache/data"
|
||||
TEMP_DIR: "./cache/cache_data"
|
||||
#
|
||||
FREEZE:
|
||||
#
|
||||
@@ -249,7 +250,7 @@ SOLVER:
|
||||
#
|
||||
OPTIMIZER:
|
||||
NAME: AdamW
|
||||
LEARNING_RATE: 0.064
|
||||
LEARNING_RATE: 0.0001
|
||||
BETAS: [ 0.9, 0.999 ]
|
||||
EPS: 1e-8
|
||||
WEIGHT_DECAY: 1e-2
|
||||
|
||||
@@ -80,13 +80,14 @@ SOLVER:
|
||||
NUM_FOLDS: 1
|
||||
ACCU_STEP: 1
|
||||
EVAL_INTERVAL: -1
|
||||
RESCALE_LR: False
|
||||
#
|
||||
WORK_DIR:
|
||||
LOG_FILE: std_log.txt
|
||||
#
|
||||
FILE_SYSTEM:
|
||||
NAME: "ModelscopeFs"
|
||||
TEMP_DIR: "./cache/data"
|
||||
TEMP_DIR: "./cache/cache_data"
|
||||
#
|
||||
FREEZE:
|
||||
#
|
||||
@@ -191,7 +192,7 @@ SOLVER:
|
||||
#
|
||||
OPTIMIZER:
|
||||
NAME: AdamW
|
||||
LEARNING_RATE: 0.064
|
||||
LEARNING_RATE: 0.0001
|
||||
BETAS: [ 0.9, 0.999 ]
|
||||
EPS: 1e-8
|
||||
WEIGHT_DECAY: 1e-2
|
||||
|
||||
@@ -33,10 +33,12 @@ class DiffusionInference():
|
||||
'''
|
||||
def __init__(self, logger=None):
|
||||
self.logger = logger
|
||||
self.is_redefine_paras = True
|
||||
self.loaded_model = {}
|
||||
self.loaded_model_name = [
|
||||
'diffusion_model', 'first_stage_model', 'cond_stage_model'
|
||||
]
|
||||
self.diffusion_insclass = GaussianDiffusion
|
||||
self.tuner_infer = TunerInference(self.logger)
|
||||
self.control_infer = ControlInference(self.logger)
|
||||
|
||||
@@ -45,7 +47,8 @@ class DiffusionInference():
|
||||
self.is_default = cfg.get('IS_DEFAULT', False)
|
||||
module_paras = self.load_default(cfg.get('DEFAULT_PARAS', None))
|
||||
assert cfg.have('MODEL')
|
||||
cfg.MODEL = self.redefine_paras(cfg.MODEL)
|
||||
if self.is_redefine_paras:
|
||||
cfg.MODEL = self.redefine_paras(cfg.MODEL)
|
||||
self.diffusion = self.load_schedule(cfg.MODEL.SCHEDULE)
|
||||
self.diffusion_model = self.infer_model(
|
||||
cfg.MODEL.DIFFUSION_MODEL, module_paras.get(
|
||||
@@ -316,8 +319,8 @@ class DiffusionInference():
|
||||
def load_schedule(self, cfg):
|
||||
parameterization = cfg.get('PARAMETERIZATION', 'eps')
|
||||
assert parameterization in [
|
||||
'eps', 'x0', 'v'
|
||||
], 'currently only supporting "eps" and "x0" and "v"'
|
||||
'eps', 'x0', 'v', 'rf'
|
||||
], 'currently only supporting "eps" and "x0" and "v" and "rf"'
|
||||
num_timesteps = cfg.get('TIMESTEPS', 1000)
|
||||
|
||||
schedule_args = {
|
||||
@@ -336,8 +339,8 @@ class DiffusionInference():
|
||||
n=num_timesteps,
|
||||
zero_terminal_snr=zero_terminal_snr,
|
||||
**schedule_args)
|
||||
diffusion = GaussianDiffusion(sigmas=sigmas,
|
||||
prediction_type=parameterization)
|
||||
diffusion = self.diffusion_insclass(sigmas=sigmas,
|
||||
prediction_type=parameterization)
|
||||
return diffusion
|
||||
|
||||
def get_batch(self, value_dict, num_samples=1):
|
||||
|
||||
@@ -8,16 +8,12 @@ from collections import OrderedDict
|
||||
import gradio as gr
|
||||
import torch
|
||||
import torchvision.transforms.functional as TF
|
||||
from scepter.modules.model.network.diffusion.diffusion import GaussianDiffusion
|
||||
from scepter.modules.model.utils.data_utils import crop_back
|
||||
from scepter.modules.utils.distribute import we
|
||||
from scepter.modules.utils.file_system import FS
|
||||
|
||||
from .diffusion_inference import DiffusionInference
|
||||
|
||||
|
||||
def get_model(model_tuple):
|
||||
assert 'model' in model_tuple
|
||||
return model_tuple['model']
|
||||
from .diffusion_inference import DiffusionInference, get_model
|
||||
|
||||
|
||||
class LargenInference(DiffusionInference):
|
||||
@@ -28,10 +24,12 @@ class LargenInference(DiffusionInference):
|
||||
'''
|
||||
def __init__(self, logger=None):
|
||||
self.logger = logger
|
||||
self.is_redefine_paras = True
|
||||
self.loaded_model = {}
|
||||
self.loaded_model_name = [
|
||||
'diffusion_model', 'first_stage_model', 'cond_stage_model'
|
||||
]
|
||||
self.diffusion_insclass = GaussianDiffusion
|
||||
|
||||
def redefine_paras(self, cfg):
|
||||
if cfg.get('PRETRAINED_MODEL', None):
|
||||
|
||||
@@ -0,0 +1,209 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import copy
|
||||
import os.path
|
||||
import random
|
||||
from collections import OrderedDict
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from PIL.Image import Image
|
||||
from scepter.modules.model.network.diffusion.diffusion import GaussianDiffusion
|
||||
from scepter.modules.model.network.diffusion.schedules import noise_schedule
|
||||
from scepter.modules.model.registry import (BACKBONES, EMBEDDERS, MODELS,
|
||||
TOKENIZERS)
|
||||
from scepter.modules.utils.distribute import we
|
||||
from scepter.modules.utils.file_system import FS
|
||||
from scepter.studio.utils.env import get_available_memory
|
||||
|
||||
from .control_inference import ControlInference
|
||||
from .diffusion_inference import DiffusionInference, get_model
|
||||
from .tuner_inference import TunerInference
|
||||
|
||||
|
||||
class PixArtInference(DiffusionInference):
|
||||
'''
|
||||
define vae, unet, text-encoder, tuner, refiner components
|
||||
support to load the components dynamicly.
|
||||
create and load model when run this model at the first time.
|
||||
'''
|
||||
def __init__(self, logger=None):
|
||||
self.logger = logger
|
||||
self.is_redefine_paras = False
|
||||
self.loaded_model = {}
|
||||
self.loaded_model_name = [
|
||||
'diffusion_model', 'first_stage_model', 'cond_stage_model'
|
||||
]
|
||||
self.diffusion_insclass = GaussianDiffusion
|
||||
self.tuner_infer = TunerInference(self.logger)
|
||||
self.control_infer = ControlInference(self.logger)
|
||||
|
||||
@torch.no_grad()
|
||||
def __call__(self,
|
||||
input,
|
||||
num_samples=1,
|
||||
intermediate_callback=None,
|
||||
refine_strength=0,
|
||||
img_to_img_strength=0,
|
||||
cat_uc=True,
|
||||
tuner_model=None,
|
||||
control_model=None,
|
||||
**kwargs):
|
||||
|
||||
value_input = copy.deepcopy(self.input)
|
||||
value_input.update(input)
|
||||
print(value_input)
|
||||
height, width = value_input['target_size_as_tuple']
|
||||
value_output = copy.deepcopy(self.output)
|
||||
|
||||
# register tuner
|
||||
if tuner_model is not None and tuner_model != '' and len(
|
||||
tuner_model) > 0:
|
||||
if not isinstance(tuner_model, list):
|
||||
tuner_model = [tuner_model]
|
||||
self.dynamic_load(self.diffusion_model, 'diffusion_model')
|
||||
self.dynamic_load(self.cond_stage_model, 'cond_stage_model')
|
||||
self.tuner_infer.register_tuner(tuner_model, self.diffusion_model,
|
||||
self.cond_stage_model)
|
||||
self.dynamic_unload(self.diffusion_model,
|
||||
'diffusion_model',
|
||||
skip_loaded=True)
|
||||
self.dynamic_unload(self.cond_stage_model,
|
||||
'cond_stage_model',
|
||||
skip_loaded=True)
|
||||
|
||||
# cond stage
|
||||
self.dynamic_load(self.cond_stage_model, 'cond_stage_model')
|
||||
function_name, dtype = self.get_function_info(self.cond_stage_model)
|
||||
with torch.autocast('cuda',
|
||||
enabled=dtype == 'float16',
|
||||
dtype=getattr(torch, dtype)):
|
||||
context, null_context = {}, {}
|
||||
cont_mask = None
|
||||
cont, cont_mask = getattr(get_model(self.cond_stage_model),
|
||||
function_name)(value_input['prompt'],
|
||||
return_mask=True)
|
||||
context['crossattn'] = cont.float()
|
||||
null_context['crossattn'] = get_model(
|
||||
self.diffusion_model).y_embedder.y_embedding[None].repeat(
|
||||
num_samples, 1, 1)
|
||||
self.dynamic_unload(self.cond_stage_model,
|
||||
'cond_stage_model',
|
||||
skip_loaded=True)
|
||||
|
||||
# get noise
|
||||
seed = kwargs.pop('seed', -1)
|
||||
g = torch.Generator(device=we.device_id)
|
||||
seed = seed if seed >= 0 else random.randint(0, 2**32 - 1)
|
||||
g.manual_seed(seed)
|
||||
if 'seed' in value_output:
|
||||
value_output['seed'] = seed
|
||||
for sample_id in range(num_samples):
|
||||
if self.diffusion_model is not None:
|
||||
noise = torch.empty(
|
||||
1,
|
||||
4,
|
||||
height // self.first_stage_model['paras']['size_factor'],
|
||||
width // self.first_stage_model['paras']['size_factor'],
|
||||
device=we.device_id).normal_(generator=g)
|
||||
|
||||
self.dynamic_load(self.diffusion_model, 'diffusion_model')
|
||||
# UNet use input n_prompt
|
||||
function_name, dtype = self.get_function_info(
|
||||
self.diffusion_model)
|
||||
with torch.autocast('cuda',
|
||||
enabled=dtype == 'float16',
|
||||
dtype=getattr(torch, dtype)):
|
||||
latent = self.diffusion.sample(
|
||||
solver=value_input.get('sample', 'ddim'),
|
||||
noise=noise,
|
||||
model=get_model(self.diffusion_model),
|
||||
model_kwargs=[{
|
||||
'cond': context,
|
||||
'mask': cont_mask,
|
||||
'data_info': {
|
||||
'img_hw':
|
||||
torch.tensor([[height, width]],
|
||||
dtype=torch.float,
|
||||
device=we.device_id).repeat(
|
||||
num_samples, 1),
|
||||
'aspect_ratio':
|
||||
torch.tensor([[1.]],
|
||||
device=we.device_id).repeat(
|
||||
num_samples, 1)
|
||||
}
|
||||
}, {
|
||||
'cond': null_context,
|
||||
'mask': cont_mask,
|
||||
'data_info': {
|
||||
'img_hw':
|
||||
torch.tensor([[height, width]],
|
||||
dtype=torch.float,
|
||||
device=we.device_id).repeat(
|
||||
num_samples, 1),
|
||||
'aspect_ratio':
|
||||
torch.tensor([[1.]],
|
||||
device=we.device_id).repeat(
|
||||
num_samples, 1)
|
||||
}
|
||||
}],
|
||||
cat_uc=False,
|
||||
steps=value_input.get('sample_steps', 50),
|
||||
guide_scale=value_input.get('guide_scale', 7.5),
|
||||
guide_rescale=value_input.get('guide_rescale', 0.5),
|
||||
discretization=value_input.get('discretization',
|
||||
'trailing'),
|
||||
show_progress=True,
|
||||
seed=seed,
|
||||
condition_fn=None,
|
||||
clamp=None,
|
||||
percentile=None,
|
||||
t_max=None,
|
||||
t_min=None,
|
||||
discard_penultimate_step=None,
|
||||
return_intermediate=None,
|
||||
**kwargs)
|
||||
|
||||
self.dynamic_unload(self.diffusion_model,
|
||||
'diffusion_model',
|
||||
skip_loaded=True)
|
||||
|
||||
if 'latent' in value_output:
|
||||
if value_output['latent'] is None or (
|
||||
isinstance(value_output['latent'], list)
|
||||
and len(value_output['latent']) < 1):
|
||||
value_output['latent'] = []
|
||||
value_output['latent'].append(latent)
|
||||
|
||||
self.dynamic_load(self.first_stage_model, 'first_stage_model')
|
||||
x_samples = self.decode_first_stage(latent).float()
|
||||
self.dynamic_unload(self.first_stage_model,
|
||||
'first_stage_model',
|
||||
skip_loaded=True)
|
||||
images = torch.clamp((x_samples + 1.0) / 2.0, min=0.0, max=1.0)
|
||||
if 'images' in value_output:
|
||||
if value_output['images'] is None or (
|
||||
isinstance(value_output['images'], list)
|
||||
and len(value_output['images']) < 1):
|
||||
value_output['images'] = []
|
||||
value_output['images'].append(images)
|
||||
|
||||
for k, v in value_output.items():
|
||||
if isinstance(v, list):
|
||||
value_output[k] = torch.cat(v, dim=0)
|
||||
if isinstance(v, torch.Tensor):
|
||||
value_output[k] = v.cpu()
|
||||
|
||||
# unregister tuner
|
||||
if tuner_model is not None and tuner_model != '' and len(
|
||||
tuner_model) > 0:
|
||||
self.tuner_infer.unregister_tuner(tuner_model,
|
||||
self.diffusion_model,
|
||||
self.cond_stage_model)
|
||||
|
||||
# unregister control
|
||||
if control_model is not None and control_model != '':
|
||||
self.control_infer.unregister_controllers(control_model,
|
||||
self.diffusion_model)
|
||||
|
||||
return value_output
|
||||
@@ -0,0 +1,189 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import copy
|
||||
import random
|
||||
|
||||
import gradio as gr
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
import torchvision.transforms.functional as TF
|
||||
from scepter.modules.model.network.diffusion.diffusion import \
|
||||
GaussianDiffusionRF
|
||||
from scepter.modules.utils.distribute import we
|
||||
|
||||
from .control_inference import ControlInference
|
||||
from .diffusion_inference import DiffusionInference, get_model
|
||||
from .tuner_inference import TunerInference
|
||||
|
||||
|
||||
class SD3Inference(DiffusionInference):
|
||||
def __init__(self, logger=None):
|
||||
self.logger = logger
|
||||
self.is_redefine_paras = False
|
||||
self.loaded_model = {}
|
||||
self.loaded_model_name = [
|
||||
'diffusion_model', 'first_stage_model', 'cond_stage_model'
|
||||
]
|
||||
self.diffusion_insclass = GaussianDiffusionRF
|
||||
self.tuner_infer = TunerInference(self.logger)
|
||||
self.control_infer = ControlInference(self.logger)
|
||||
|
||||
@torch.no_grad()
|
||||
def decode_first_stage(self, z):
|
||||
_, dtype = self.get_function_info(self.first_stage_model, 'decode')
|
||||
# with torch.autocast('cuda', enabled=dtype == 'float16', dtype=getattr(torch, dtype)):
|
||||
z = 1. / self.first_stage_model['paras'][
|
||||
'scale_factor'] * z + self.first_stage_model['paras'][
|
||||
'shift_factor']
|
||||
return get_model(self.first_stage_model).decode(z)
|
||||
|
||||
@torch.no_grad()
|
||||
def __call__(self,
|
||||
input,
|
||||
num_samples=1,
|
||||
cat_uc=True,
|
||||
tuner_model=None,
|
||||
control_model=None,
|
||||
**kwargs):
|
||||
|
||||
value_input = copy.deepcopy(self.input)
|
||||
value_input.update(input)
|
||||
print(value_input)
|
||||
height, width = value_input['target_size_as_tuple']
|
||||
value_output = copy.deepcopy(self.output)
|
||||
|
||||
# register tuner
|
||||
if tuner_model is not None and tuner_model != '' and len(
|
||||
tuner_model) > 0:
|
||||
if not isinstance(tuner_model, list):
|
||||
tuner_model = [tuner_model]
|
||||
self.dynamic_load(self.diffusion_model, 'diffusion_model')
|
||||
self.dynamic_load(self.cond_stage_model, 'cond_stage_model')
|
||||
self.tuner_infer.register_tuner(tuner_model, self.diffusion_model,
|
||||
self.cond_stage_model)
|
||||
self.dynamic_unload(self.diffusion_model,
|
||||
'diffusion_model',
|
||||
skip_loaded=True)
|
||||
self.dynamic_unload(self.cond_stage_model,
|
||||
'cond_stage_model',
|
||||
skip_loaded=True)
|
||||
|
||||
# cond stage
|
||||
self.dynamic_load(self.cond_stage_model, 'cond_stage_model')
|
||||
function_name, dtype = self.get_function_info(self.cond_stage_model)
|
||||
context, null_context = {}, {}
|
||||
with torch.autocast('cuda',
|
||||
enabled=dtype == 'float16',
|
||||
dtype=getattr(torch, dtype)):
|
||||
ctx, pooled = getattr(get_model(self.cond_stage_model),
|
||||
function_name)(value_input['prompt'])
|
||||
null_ctx, null_pooled = getattr(get_model(self.cond_stage_model),
|
||||
function_name)([''] * num_samples)
|
||||
context['crossattn'] = ctx.float()
|
||||
context['y'] = pooled.float()
|
||||
null_context['crossattn'] = null_ctx.float()
|
||||
null_context['y'] = null_pooled.float()
|
||||
self.dynamic_unload(self.cond_stage_model,
|
||||
'cond_stage_model',
|
||||
skip_loaded=True)
|
||||
|
||||
# get noise
|
||||
seed = kwargs.pop('seed', -1)
|
||||
g = torch.Generator(device=we.device_id)
|
||||
seed = seed if seed >= 0 else random.randint(0, 2**32 - 1)
|
||||
g.manual_seed(seed)
|
||||
if 'seed' in value_output:
|
||||
value_output['seed'] = seed
|
||||
for sample_id in range(num_samples):
|
||||
if self.diffusion_model is not None:
|
||||
noise = torch.empty(
|
||||
1,
|
||||
16,
|
||||
height // self.first_stage_model['paras']['size_factor'],
|
||||
width // self.first_stage_model['paras']['size_factor'],
|
||||
device=we.device_id).normal_(generator=g)
|
||||
|
||||
self.dynamic_load(self.diffusion_model, 'diffusion_model')
|
||||
# UNet use input n_prompt
|
||||
function_name, dtype = self.get_function_info(
|
||||
self.diffusion_model)
|
||||
with torch.autocast('cuda',
|
||||
enabled=dtype == 'float16',
|
||||
dtype=getattr(torch, dtype)):
|
||||
solver_sample = value_input.get('sample', 'ddim')
|
||||
sample_steps = value_input.get('sample_steps', 50)
|
||||
guide_scale = value_input.get('guide_scale', 7.5)
|
||||
guide_rescale = value_input.get('guide_rescale', 0.5)
|
||||
discretization = value_input.get('discretization',
|
||||
'trailing')
|
||||
latent = self.diffusion.sample(
|
||||
solver=solver_sample,
|
||||
noise=noise,
|
||||
model=get_model(self.diffusion_model),
|
||||
model_kwargs=[{
|
||||
'cond': context
|
||||
}, {
|
||||
'cond': null_context
|
||||
}]
|
||||
if guide_scale is not None and guide_scale > 0 else {
|
||||
'cond': context,
|
||||
},
|
||||
cat_uc=False,
|
||||
steps=sample_steps,
|
||||
guide_scale=guide_scale,
|
||||
guide_rescale=guide_rescale,
|
||||
discretization=discretization,
|
||||
show_progress=True,
|
||||
seed=seed,
|
||||
condition_fn=None,
|
||||
clamp=None,
|
||||
percentile=None,
|
||||
t_max=None,
|
||||
t_min=None,
|
||||
discard_penultimate_step=None,
|
||||
return_intermediate=None,
|
||||
**kwargs)
|
||||
|
||||
self.dynamic_unload(self.diffusion_model,
|
||||
'diffusion_model',
|
||||
skip_loaded=True)
|
||||
|
||||
if 'latent' in value_output:
|
||||
if value_output['latent'] is None or (
|
||||
isinstance(value_output['latent'], list)
|
||||
and len(value_output['latent']) < 1):
|
||||
value_output['latent'] = []
|
||||
value_output['latent'].append(latent)
|
||||
|
||||
self.dynamic_load(self.first_stage_model, 'first_stage_model')
|
||||
x_samples = self.decode_first_stage(latent).float()
|
||||
self.dynamic_unload(self.first_stage_model,
|
||||
'first_stage_model',
|
||||
skip_loaded=True)
|
||||
images = torch.clamp((x_samples + 1.0) / 2.0, min=0.0, max=1.0)
|
||||
if 'images' in value_output:
|
||||
if value_output['images'] is None or (
|
||||
isinstance(value_output['images'], list)
|
||||
and len(value_output['images']) < 1):
|
||||
value_output['images'] = []
|
||||
value_output['images'].append(images)
|
||||
|
||||
for k, v in value_output.items():
|
||||
if isinstance(v, list):
|
||||
value_output[k] = torch.cat(v, dim=0)
|
||||
if isinstance(v, torch.Tensor):
|
||||
value_output[k] = v.cpu()
|
||||
|
||||
# unregister tuner
|
||||
if tuner_model is not None and tuner_model != '' and len(
|
||||
tuner_model) > 0:
|
||||
self.tuner_infer.unregister_tuner(tuner_model,
|
||||
self.diffusion_model,
|
||||
self.cond_stage_model)
|
||||
|
||||
# unregister control
|
||||
if control_model is not None and control_model != '':
|
||||
self.control_infer.unregister_controllers(control_model,
|
||||
self.diffusion_model)
|
||||
|
||||
return value_output
|
||||
@@ -7,18 +7,14 @@ import gradio as gr
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
import torchvision.transforms.functional as TF
|
||||
from scepter.modules.model.network.diffusion.diffusion import GaussianDiffusion
|
||||
from scepter.modules.utils.distribute import we
|
||||
|
||||
from .control_inference import ControlInference
|
||||
from .diffusion_inference import DiffusionInference
|
||||
from .diffusion_inference import DiffusionInference, get_model
|
||||
from .tuner_inference import TunerInference
|
||||
|
||||
|
||||
def get_model(model_tuple):
|
||||
assert 'model' in model_tuple
|
||||
return model_tuple['model']
|
||||
|
||||
|
||||
class StyleboothInference(DiffusionInference):
|
||||
'''
|
||||
define vae, unet, text-encoder, tuner, refiner components
|
||||
@@ -27,10 +23,12 @@ class StyleboothInference(DiffusionInference):
|
||||
'''
|
||||
def __init__(self, logger=None):
|
||||
self.logger = logger
|
||||
self.is_redefine_paras = True
|
||||
self.loaded_model = {}
|
||||
self.loaded_model_name = [
|
||||
'diffusion_model', 'first_stage_model', 'cond_stage_model'
|
||||
]
|
||||
self.diffusion_insclass = GaussianDiffusion
|
||||
self.tuner_infer = TunerInference(self.logger)
|
||||
self.control_infer = ControlInference(self.logger)
|
||||
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
from scepter.modules.model.backbone import (autoencoder, image, unet, utils,
|
||||
video)
|
||||
from scepter.modules.model.backbone import (autoencoder, image, mmdit, pixart,
|
||||
unet, utils, video)
|
||||
|
||||
@@ -0,0 +1,2 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
from .sd3 import MMDiT
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,2 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
from .pixart_alpha import PixArt
|
||||
@@ -0,0 +1,503 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
# All rights reserved.
|
||||
# This file contains code that is adapted from
|
||||
# timm: https://github.com/huggingface/pytorch-image-models
|
||||
# pixart: https://github.com/PixArt-alpha/PixArt-alpha
|
||||
|
||||
# This source code is also licensed under the license found in the
|
||||
# LICENSE file in the root directory of this source tree.
|
||||
# --------------------------------------------------------
|
||||
|
||||
from collections import OrderedDict
|
||||
# MAE: https://github.com/facebookresearch/mae/blob/main/models_mae.py
|
||||
# --------------------------------------------------------
|
||||
from typing import Iterable
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
# References:
|
||||
# GLIDE: https://github.com/openai/glide-text2im
|
||||
from torch.utils.checkpoint import checkpoint, checkpoint_sequential
|
||||
|
||||
from scepter.modules.model.backbone.transformer.attention import \
|
||||
MultiHeadAttention
|
||||
from scepter.modules.model.backbone.transformer.layers import (
|
||||
CaptionEmbedder, DropPath, LabelEmbedder, Mlp, SizeEmbedder,
|
||||
TimestepEmbedder, modulate)
|
||||
from scepter.modules.model.backbone.transformer.patchify import (PatchEmbed,
|
||||
unpatchify)
|
||||
from scepter.modules.model.backbone.transformer.pos_embed import \
|
||||
get_2d_sincos_pos_embed
|
||||
from scepter.modules.model.base_model import BaseModel
|
||||
from scepter.modules.model.registry import BACKBONES
|
||||
from scepter.modules.utils.config import dict_to_yaml
|
||||
from scepter.modules.utils.file_system import FS
|
||||
|
||||
|
||||
def auto_grad_checkpoint(module, *args, use_grad_checkpoint=False, **kwargs):
|
||||
if use_grad_checkpoint:
|
||||
if not isinstance(module, Iterable):
|
||||
return checkpoint(module, *args, use_reentrant=False, **kwargs)
|
||||
gc_step = module[0].grad_checkpointing_step
|
||||
return checkpoint_sequential(module,
|
||||
gc_step,
|
||||
*args,
|
||||
use_reentrant=False,
|
||||
**kwargs)
|
||||
return module(*args, **kwargs)
|
||||
|
||||
|
||||
class FinalLayer(nn.Module):
|
||||
"""
|
||||
The final layer of PixArt.
|
||||
"""
|
||||
def __init__(self, hidden_size, patch_size, out_channels):
|
||||
super().__init__()
|
||||
self.norm_final = nn.LayerNorm(hidden_size,
|
||||
elementwise_affine=False,
|
||||
eps=1e-6)
|
||||
self.linear = nn.Linear(hidden_size,
|
||||
patch_size * patch_size * out_channels,
|
||||
bias=True)
|
||||
self.adaLN_modulation = nn.Sequential(
|
||||
nn.SiLU(), nn.Linear(hidden_size, 2 * hidden_size, bias=True))
|
||||
|
||||
def forward(self, x, c):
|
||||
shift, scale = self.adaLN_modulation(c).chunk(2, dim=1)
|
||||
# modulate x
|
||||
x = modulate(self.norm_final(x), shift, scale, unsqueeze=True)
|
||||
x = self.linear(x)
|
||||
return x
|
||||
|
||||
|
||||
class T2IFinalLayer(nn.Module):
|
||||
"""
|
||||
The final layer of PixArt.
|
||||
"""
|
||||
def __init__(self, hidden_size, patch_size, out_channels):
|
||||
super().__init__()
|
||||
self.norm_final = nn.LayerNorm(hidden_size,
|
||||
elementwise_affine=False,
|
||||
eps=1e-6)
|
||||
self.linear = nn.Linear(hidden_size,
|
||||
patch_size * patch_size * out_channels,
|
||||
bias=True)
|
||||
self.scale_shift_table = nn.Parameter(
|
||||
torch.randn(2, hidden_size) / hidden_size**0.5)
|
||||
self.out_channels = out_channels
|
||||
|
||||
def forward(self, x, t):
|
||||
shift, scale = (self.scale_shift_table[None] + t[:, None]).chunk(2,
|
||||
dim=1)
|
||||
x = modulate(self.norm_final(x), shift, scale)
|
||||
x = self.linear(x)
|
||||
return x
|
||||
|
||||
|
||||
class DitFinalLayer(nn.Module):
|
||||
"""
|
||||
The final layer of PixArt.
|
||||
"""
|
||||
def __init__(self, hidden_size, patch_size, out_channels):
|
||||
super().__init__()
|
||||
self.norm_final = nn.LayerNorm(hidden_size,
|
||||
elementwise_affine=False,
|
||||
eps=1e-6)
|
||||
self.linear = nn.Linear(hidden_size,
|
||||
patch_size * patch_size * out_channels,
|
||||
bias=True)
|
||||
self.adaLN_modulation = nn.Sequential(
|
||||
nn.SiLU(), nn.Linear(hidden_size, 2 * hidden_size, bias=True))
|
||||
self.out_channels = out_channels
|
||||
|
||||
def forward(self, x, t):
|
||||
shift, scale = self.adaLN_modulation(t).chunk(2, dim=1)
|
||||
x = modulate(self.norm_final(x), shift, scale, unsqueeze=True)
|
||||
x = self.linear(x)
|
||||
return x
|
||||
|
||||
|
||||
class PixArtBlock(nn.Module):
|
||||
"""
|
||||
A PixArt block with adaptive layer norm zero (adaLN-Zero) conditioning.
|
||||
"""
|
||||
def __init__(self,
|
||||
hidden_size,
|
||||
num_heads,
|
||||
mlp_ratio=4.0,
|
||||
drop_path=0.,
|
||||
window_size=0,
|
||||
use_rel_pos=False,
|
||||
backend=None,
|
||||
use_condition=True,
|
||||
**block_kwargs):
|
||||
super().__init__()
|
||||
self.hidden_size = hidden_size
|
||||
self.use_condition = use_condition
|
||||
self.norm1 = nn.LayerNorm(hidden_size,
|
||||
elementwise_affine=False,
|
||||
eps=1e-6)
|
||||
self.attn = MultiHeadAttention(hidden_size,
|
||||
num_heads=num_heads,
|
||||
qkv_bias=True,
|
||||
backend=backend,
|
||||
**block_kwargs)
|
||||
if self.use_condition:
|
||||
self.cross_attn = MultiHeadAttention(hidden_size,
|
||||
context_dim=hidden_size,
|
||||
num_heads=num_heads,
|
||||
qkv_bias=True,
|
||||
backend=backend,
|
||||
**block_kwargs)
|
||||
self.norm2 = nn.LayerNorm(hidden_size,
|
||||
elementwise_affine=False,
|
||||
eps=1e-6)
|
||||
# to be compatible with lower version pytorch
|
||||
approx_gelu = lambda: nn.GELU(approximate='tanh')
|
||||
self.mlp = Mlp(in_features=hidden_size,
|
||||
hidden_features=int(hidden_size * mlp_ratio),
|
||||
act_layer=approx_gelu,
|
||||
drop=0)
|
||||
self.drop_path = DropPath(
|
||||
drop_path) if drop_path > 0. else nn.Identity()
|
||||
self.window_size = window_size
|
||||
self.scale_shift_table = nn.Parameter(
|
||||
torch.randn(6, hidden_size) / hidden_size**0.5)
|
||||
|
||||
def forward(self, x, y, t, mask=None, **kwargs):
|
||||
B, N, C = x.shape
|
||||
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = (
|
||||
self.scale_shift_table[None] + t.reshape(B, 6, -1)).chunk(6, dim=1)
|
||||
x = x + self.drop_path(gate_msa * self.attn(
|
||||
modulate(self.norm1(x), shift_msa, scale_msa, unsqueeze=False)))
|
||||
if self.use_condition:
|
||||
x = x + self.cross_attn(x, y, mask)
|
||||
x = x + self.drop_path(gate_mlp * self.mlp(
|
||||
modulate(self.norm2(x), shift_mlp, scale_mlp, unsqueeze=False)))
|
||||
return x
|
||||
|
||||
|
||||
@BACKBONES.register_class()
|
||||
class PixArt(BaseModel):
|
||||
"""
|
||||
Diffusion model with a Transformer backbone.
|
||||
"""
|
||||
para_dict = BaseModel.para_dict
|
||||
para_dict.update({
|
||||
'PATCH_SIZE': {
|
||||
'value': 2,
|
||||
'description': ''
|
||||
},
|
||||
'IN_CHANNELS': {
|
||||
'value': 4,
|
||||
'description': ''
|
||||
},
|
||||
'HIDDEN_SIZE': {
|
||||
'value': 1152,
|
||||
'description': ''
|
||||
},
|
||||
'DEPTH': {
|
||||
'value': 28,
|
||||
'description': ''
|
||||
},
|
||||
'NUM_HEADS': {
|
||||
'value': 16,
|
||||
'description': ''
|
||||
},
|
||||
'MLP_RATIO': {
|
||||
'value': 4.0,
|
||||
'description': ''
|
||||
},
|
||||
'CLASS_DROPOUT_PROB': {
|
||||
'value': 0.1,
|
||||
'description': ''
|
||||
},
|
||||
'PRED_SIGMA': {
|
||||
'value': True,
|
||||
'description': ''
|
||||
},
|
||||
'DROP_PATH': {
|
||||
'value': 0.,
|
||||
'description': ''
|
||||
},
|
||||
'WINDOW_DIZE': {
|
||||
'value': 0,
|
||||
'description': ''
|
||||
},
|
||||
'WINDOW_BLOCK_INDEXES': {
|
||||
'value': None,
|
||||
'description': ''
|
||||
},
|
||||
'USE_REL_POS': {
|
||||
'value': False,
|
||||
'description': ''
|
||||
},
|
||||
'CAPTION_CHANNELS': {
|
||||
'value': 4096,
|
||||
'description': ''
|
||||
},
|
||||
'USE_AR_SIZE': {
|
||||
'value': True,
|
||||
'description': ''
|
||||
},
|
||||
'DIT_FINAL_LAYER': {
|
||||
'value': False,
|
||||
'description': ''
|
||||
},
|
||||
'LEWEI_SCALE': {
|
||||
'value': 1.0,
|
||||
'description': ''
|
||||
},
|
||||
'MODEL_MAX_LENGTH': {
|
||||
'value': 120,
|
||||
'description': ''
|
||||
},
|
||||
'NUM_CLASSES': {
|
||||
'value':
|
||||
None,
|
||||
'description':
|
||||
'The class num for class guided setting, also can be set as continuous.'
|
||||
},
|
||||
'ATTENTION_BACKEND': {
|
||||
'value': None,
|
||||
'description': ''
|
||||
}
|
||||
})
|
||||
|
||||
def __init__(self, cfg, logger):
|
||||
super().__init__(cfg, logger=logger)
|
||||
self.window_block_indexes = cfg.get('WINDOW_BLOCK_INDEXES', None)
|
||||
if self.window_block_indexes is None:
|
||||
self.window_block_indexes = []
|
||||
self.pred_sigma = cfg.get('PRED_SIGMA', True)
|
||||
self.in_channels = cfg.get('IN_CHANNELS', 4)
|
||||
self.out_channels = self.in_channels * 2 if self.pred_sigma else self.in_channels
|
||||
self.patch_size = cfg.get('PATCH_SIZE', 2)
|
||||
self.num_heads = cfg.get('NUM_HEADS', 16)
|
||||
self.hidden_size = cfg.get('HIDDEN_SIZE', 1152)
|
||||
self.lewei_scale = cfg.get('LEWEI_SCALE', 1.0),
|
||||
self.caption_channels = cfg.get('CAPTION_CHANNELS', 4096)
|
||||
self.class_dropout_prob = cfg.get('CLASS_DROPOUT_PROB', 0.1)
|
||||
self.model_max_length = cfg.get('MODEL_MAX_LENGTH', 120)
|
||||
self.drop_path = cfg.get('DROP_PATH', 0.)
|
||||
self.depth = cfg.get('DEPTH', 28)
|
||||
self.mlp_ratio = cfg.get('MLP_RATIO', 4.0)
|
||||
self.num_classes = cfg.get('NUM_CLASSES', None)
|
||||
self.use_grad_checkpoint = cfg.get('USE_GRAD_CHECKPOINT', False)
|
||||
self.use_ar_size = cfg.get('USE_AR_SIZE', True)
|
||||
self.use_dit_final_layer = cfg.get('DIT_FINAL_LAYER', False)
|
||||
self.attention_backend = cfg.get('ATTENTION_BACKEND', None)
|
||||
self.ignore_keys = cfg.get('IGNORE_KEYS', [])
|
||||
|
||||
if self.num_classes is not None:
|
||||
if isinstance(self.num_classes, int):
|
||||
self.label_embedder = LabelEmbedder(
|
||||
self.num_classes,
|
||||
self.hidden_size,
|
||||
dropout_prob=self.class_dropout_prob)
|
||||
elif self.num_classes == 'continuous':
|
||||
print('setting up linear c_adm embedding layer')
|
||||
self.label_embedder = nn.Linear(1, self.hidden_size)
|
||||
else:
|
||||
raise ValueError()
|
||||
|
||||
self.x_embedder = PatchEmbed(self.patch_size,
|
||||
self.in_channels,
|
||||
self.hidden_size,
|
||||
bias=True)
|
||||
self.t_embedder = TimestepEmbedder(self.hidden_size)
|
||||
|
||||
# self.base_size = self.input_size // self.patch_size
|
||||
approx_gelu = lambda: nn.GELU(approximate='tanh')
|
||||
self.t_block = nn.Sequential(
|
||||
nn.SiLU(),
|
||||
nn.Linear(self.hidden_size, 6 * self.hidden_size, bias=True))
|
||||
if self.num_classes is None:
|
||||
self.y_embedder = CaptionEmbedder(
|
||||
in_channels=self.caption_channels,
|
||||
hidden_size=self.hidden_size,
|
||||
uncond_prob=self.class_dropout_prob,
|
||||
act_layer=approx_gelu,
|
||||
token_num=self.model_max_length)
|
||||
if self.use_ar_size:
|
||||
self.csize_embedder = SizeEmbedder(self.hidden_size //
|
||||
3) # c_size embed
|
||||
self.ar_embedder = SizeEmbedder(self.hidden_size //
|
||||
3) # aspect ratio embed
|
||||
|
||||
drop_path = [
|
||||
x.item() for x in torch.linspace(0, self.drop_path, self.depth)
|
||||
] # stochastic depth decay rule
|
||||
self.blocks = nn.ModuleList([
|
||||
PixArtBlock(self.hidden_size,
|
||||
self.num_heads,
|
||||
mlp_ratio=self.mlp_ratio,
|
||||
drop_path=drop_path[i],
|
||||
window_size=self.window_size
|
||||
if i in self.window_block_indexes else 0,
|
||||
use_rel_pos=self.use_rel_pos
|
||||
if i in self.window_block_indexes else False,
|
||||
backend=self.attention_backend,
|
||||
use_condition=self.num_classes is None)
|
||||
for i in range(self.depth)
|
||||
])
|
||||
if self.use_dit_final_layer:
|
||||
self.final_layer = DitFinalLayer(self.hidden_size, self.patch_size,
|
||||
self.out_channels)
|
||||
else:
|
||||
self.final_layer = T2IFinalLayer(self.hidden_size, self.patch_size,
|
||||
self.out_channels)
|
||||
|
||||
self.initialize_weights()
|
||||
|
||||
def load_pretrained_model(self, pretrained_model):
|
||||
if pretrained_model:
|
||||
with FS.get_from(pretrained_model, wait_finish=True) as local_path:
|
||||
model = torch.load(local_path, map_location='cpu')
|
||||
if 'state_dict' in model:
|
||||
model = model['state_dict']
|
||||
new_ckpt = OrderedDict()
|
||||
for k, v in model.items():
|
||||
k = k.replace('.cross_attn.q_linear.', '.cross_attn.q.')
|
||||
k = k.replace('.cross_attn.proj.',
|
||||
'.cross_attn.o.').replace(
|
||||
'.attn.proj.', '.attn.o.')
|
||||
if '.cross_attn.kv_linear.' in k:
|
||||
k_p, v_p = torch.split(v, v.shape[0] // 2)
|
||||
new_ckpt[k.replace('.cross_attn.kv_linear.',
|
||||
'.cross_attn.k.')] = k_p
|
||||
new_ckpt[k.replace('.cross_attn.kv_linear.',
|
||||
'.cross_attn.v.')] = v_p
|
||||
elif '.attn.qkv.' in k:
|
||||
q_p, k_p, v_p = torch.split(v, v.shape[0] // 3)
|
||||
new_ckpt[k.replace('.attn.qkv.', '.attn.q.')] = q_p
|
||||
new_ckpt[k.replace('.attn.qkv.', '.attn.k.')] = k_p
|
||||
new_ckpt[k.replace('.attn.qkv.', '.attn.v.')] = v_p
|
||||
else:
|
||||
new_ckpt[k] = v
|
||||
missing, unexpected = self.load_state_dict(new_ckpt,
|
||||
strict=False)
|
||||
print(
|
||||
f'Restored from {pretrained_model} with {len(missing)} missing and {len(unexpected)} unexpected keys'
|
||||
)
|
||||
if len(missing) > 0:
|
||||
print(f'Missing Keys:\n {missing}')
|
||||
if len(unexpected) > 0:
|
||||
print(f'\nUnexpected Keys:\n {unexpected}')
|
||||
|
||||
def forward(self,
|
||||
x,
|
||||
t=None,
|
||||
cond=dict(),
|
||||
mask=None,
|
||||
data_info=None,
|
||||
**kwargs):
|
||||
"""
|
||||
Forward pass of PixArt.
|
||||
x: (N, C, H, W) tensor of spatial inputs (images or latent representations of images)
|
||||
t: (N,) tensor of diffusion timesteps
|
||||
y: (N, 1, 120, C) tensor of class labels
|
||||
"""
|
||||
label = None
|
||||
if isinstance(cond, dict):
|
||||
if 'label' in cond and cond['label'] is not None:
|
||||
label = cond['label']
|
||||
if 'concat' in cond:
|
||||
concat = cond['concat']
|
||||
x = torch.cat([x, concat], dim=1)
|
||||
context = cond.get('crossattn', None)
|
||||
else:
|
||||
context = cond
|
||||
|
||||
y = context
|
||||
h, w = x.shape[-2] // self.patch_size, x.shape[-1] // self.patch_size
|
||||
|
||||
x = self.x_embedder(x) # (N, T, D), where T = H * W / patch_size ** 2
|
||||
pos_embed = torch.from_numpy(
|
||||
get_2d_sincos_pos_embed(self.hidden_size, (h, w),
|
||||
lewei_scale=self.lewei_scale,
|
||||
base_h_size=h,
|
||||
base_w_size=w)).unsqueeze(0).float().to(
|
||||
x.device)
|
||||
x = x + pos_embed
|
||||
|
||||
t = self.t_embedder(t) # (N, D)
|
||||
if self.num_classes is not None and label is not None:
|
||||
t = t + self.label_embedder(label, self.training)
|
||||
if self.use_ar_size and data_info is not None:
|
||||
bs = x.shape[0]
|
||||
c_size, ar = data_info['img_hw'], data_info['aspect_ratio']
|
||||
csize = self.csize_embedder(c_size, bs) # (N, D)
|
||||
ar = self.ar_embedder(ar, bs) # (N, D)
|
||||
t = t + torch.cat([csize, ar], dim=1)
|
||||
t0 = self.t_block(t)
|
||||
if self.num_classes is not None:
|
||||
y = None
|
||||
else:
|
||||
y = self.y_embedder(y, self.training)
|
||||
for block in self.blocks:
|
||||
x = auto_grad_checkpoint(
|
||||
block,
|
||||
x,
|
||||
y,
|
||||
t0,
|
||||
mask,
|
||||
use_grad_checkpoint=self.use_grad_checkpoint)
|
||||
# (N, T, D) #support grad checkpoint
|
||||
x = self.final_layer(x, t) # (N, T, patch_size ** 2 * out_channels)
|
||||
x = unpatchify(x, h, w, self.out_channels, self.patch_size,
|
||||
self.patch_size) # (N, out_channels, H, W)
|
||||
if self.pred_sigma:
|
||||
return x.chunk(2, dim=1)[0]
|
||||
else:
|
||||
return x
|
||||
|
||||
def initialize_weights(self):
|
||||
# Initialize transformer layers:
|
||||
def _basic_init(module):
|
||||
if isinstance(module, nn.Linear):
|
||||
torch.nn.init.xavier_uniform_(module.weight)
|
||||
if module.bias is not None:
|
||||
nn.init.constant_(module.bias, 0)
|
||||
|
||||
self.apply(_basic_init)
|
||||
|
||||
# Initialize patch_embed like nn.Linear (instead of nn.Conv2d):
|
||||
w = self.x_embedder.proj.weight.data
|
||||
nn.init.xavier_uniform_(w.view([w.shape[0], -1]))
|
||||
# Initialize timestep embedding MLP:
|
||||
nn.init.normal_(self.t_embedder.mlp[0].weight, std=0.02)
|
||||
nn.init.normal_(self.t_embedder.mlp[2].weight, std=0.02)
|
||||
nn.init.normal_(self.t_block[1].weight, std=0.02)
|
||||
if self.use_ar_size:
|
||||
nn.init.normal_(self.csize_embedder.mlp[0].weight, std=0.02)
|
||||
nn.init.normal_(self.csize_embedder.mlp[2].weight, std=0.02)
|
||||
nn.init.normal_(self.ar_embedder.mlp[0].weight, std=0.02)
|
||||
nn.init.normal_(self.ar_embedder.mlp[2].weight, std=0.02)
|
||||
if self.num_classes is not None:
|
||||
nn.init.normal_(self.label_embedder.embedding_table.weight,
|
||||
std=0.02)
|
||||
# Initialize caption embedding MLP:
|
||||
if hasattr(self, 'y_embedder'):
|
||||
nn.init.normal_(self.y_embedder.y_proj.fc1.weight, std=0.02)
|
||||
nn.init.normal_(self.y_embedder.y_proj.fc2.weight, std=0.02)
|
||||
# Zero-out adaLN modulation layers in PixArt blocks:
|
||||
if self.num_classes is None:
|
||||
for block in self.blocks:
|
||||
nn.init.constant_(block.cross_attn.o.weight, 0)
|
||||
nn.init.constant_(block.cross_attn.o.bias, 0)
|
||||
# Zero-out output layers:
|
||||
nn.init.constant_(self.final_layer.linear.weight, 0)
|
||||
nn.init.constant_(self.final_layer.linear.bias, 0)
|
||||
|
||||
@property
|
||||
def dtype(self):
|
||||
return next(self.parameters()).dtype
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
return dict_to_yaml('BACKBONE',
|
||||
__class__.__name__,
|
||||
PixArt.para_dict,
|
||||
set_name=True)
|
||||
@@ -0,0 +1,760 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
# All rights reserved.
|
||||
# This file contains code that is adapted from
|
||||
# timm: https://github.com/huggingface/pytorch-image-models
|
||||
# pixart: https://github.com/PixArt-alpha/PixArt-alpha
|
||||
import math
|
||||
import time
|
||||
import warnings
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from torch.cuda import amp
|
||||
from torch.nn import functional as F
|
||||
from torch.nn.utils.rnn import pad_sequence
|
||||
from tqdm import tqdm
|
||||
|
||||
from scepter.modules.model.backbone.transformer.pos_embed import apply_2d_rope
|
||||
|
||||
try:
|
||||
import xformers
|
||||
import xformers.ops
|
||||
XFORMERS_IS_AVAILABLE = True
|
||||
except Exception as e:
|
||||
XFORMERS_IS_AVAILABLE = False
|
||||
warnings.warn(f'{e}')
|
||||
try:
|
||||
from flash_attn import (flash_attn_varlen_func)
|
||||
FLASHATTN_IS_AVAILABLE = True
|
||||
except ImportError:
|
||||
FLASHATTN_IS_AVAILABLE = False
|
||||
flash_attn_varlen_func = None
|
||||
|
||||
|
||||
def drop_path(x, drop_prob: float = 0., training: bool = False):
|
||||
"""Drop paths (Stochastic Depth) per sample (when applied in main path of residual blocks).
|
||||
This is the same as the DropConnect impl I created for EfficientNet, etc networks, however,
|
||||
the original name is misleading as 'Drop Connect' is a different form of dropout in a separate paper...
|
||||
See discussion: https://github.com/tensorflow/tpu/issues/494#issuecomment-532968956 ... I've opted for
|
||||
changing the layer and argument names to 'drop path' rather than mix DropConnect as a layer name and use
|
||||
'survival rate' as the argument.
|
||||
"""
|
||||
if drop_prob == 0. or not training:
|
||||
return x
|
||||
keep_prob = 1 - drop_prob
|
||||
shape = (x.shape[0], ) + (1, ) * (
|
||||
x.ndim - 1) # work with diff dim tensors, not just 2D ConvNets
|
||||
random_tensor = keep_prob + torch.rand(
|
||||
shape, dtype=x.dtype, device=x.device)
|
||||
random_tensor.floor_() # binarize
|
||||
output = x.div(keep_prob) * random_tensor
|
||||
return output
|
||||
|
||||
|
||||
class MultiHeadAttention(nn.Module):
|
||||
def __init__(self,
|
||||
dim,
|
||||
context_dim=None,
|
||||
num_heads=None,
|
||||
head_dim=None,
|
||||
attn_drop=0.0,
|
||||
qkv_bias=False,
|
||||
dropout=0.0,
|
||||
backend=None,
|
||||
**block_kwargs):
|
||||
super().__init__()
|
||||
# consider head_dim first, then num_heads
|
||||
num_heads = dim // head_dim if head_dim else num_heads
|
||||
head_dim = dim // num_heads
|
||||
assert num_heads * head_dim == dim
|
||||
context_dim = context_dim or dim
|
||||
self.dim = dim
|
||||
self.context_dim = context_dim
|
||||
self.num_heads = num_heads
|
||||
self.head_dim = head_dim
|
||||
self.scale = math.pow(head_dim, -0.25)
|
||||
# layers
|
||||
self.q = nn.Linear(dim, dim, bias=qkv_bias)
|
||||
self.k = nn.Linear(context_dim, dim, bias=qkv_bias)
|
||||
self.v = nn.Linear(context_dim, dim, bias=qkv_bias)
|
||||
self.o = nn.Linear(dim, dim)
|
||||
self.dropout = nn.Dropout(dropout)
|
||||
self.attention_op = None
|
||||
self.attn_drop = nn.Dropout(attn_drop)
|
||||
self.backend = backend
|
||||
assert self.backend in ('flash_attn', 'xformer_attn', 'pytorch_attn',
|
||||
None)
|
||||
if FLASHATTN_IS_AVAILABLE and self.backend in ('flash_attn', None):
|
||||
self.backend = 'flash_attn'
|
||||
self.softmax_scale = block_kwargs.get('softmax_scale', None)
|
||||
self.causal = block_kwargs.get('causal', False)
|
||||
self.window_size = block_kwargs.get('window_size', (-1, -1))
|
||||
self.deterministic = block_kwargs.get('deterministic', False)
|
||||
elif XFORMERS_IS_AVAILABLE and self.backend in ('xformer_attn', None):
|
||||
self.backend = 'xformer_attn'
|
||||
else:
|
||||
self.backend = 'pytorch_attn'
|
||||
|
||||
def xformer_attn(self, x, context=None, mask=None, **kwargs):
|
||||
context = x if context is None else context
|
||||
b, n, d = x.size(0), self.num_heads, self.head_dim
|
||||
# compute query, key, value
|
||||
q = self.q(x).view(b, -1, n, d)
|
||||
k = self.k(context).view(b, -1, n, d)
|
||||
v = self.v(context).view(b, -1, n, d)
|
||||
|
||||
attn_bias = None
|
||||
if mask is not None:
|
||||
assert mask.ndim in [2, 3]
|
||||
mask = mask.view(b, 1, 1,
|
||||
-1) if mask.ndim == 2 else mask.unsqueeze(1)
|
||||
# To use an `attn_bias` with a sequence length that is not a multiple of 8,
|
||||
# you need to ensure memory is aligned by slicing a bigger tensor.
|
||||
# Example: use `attn_bias = torch.zeros([1, 1, 5, 8])[:,:,:,:5]`
|
||||
# instead of `torch.zeros([1, 1, 5, 5])
|
||||
q_size = math.ceil(q.size(1) / 8) * 8
|
||||
k_size = math.ceil(k.size(1) / 8) * 8
|
||||
attn_bias = x.new_zeros(b, n, q_size,
|
||||
k_size)[:, :, :q.size(1), :k.size(1)]
|
||||
attn_bias = attn_bias.masked_fill_(mask == 0,
|
||||
torch.finfo(x.dtype).min).to(
|
||||
q.dtype)
|
||||
x = xformers.ops.memory_efficient_attention(q,
|
||||
k,
|
||||
v,
|
||||
p=self.attn_drop.p,
|
||||
attn_bias=attn_bias)
|
||||
x = x.reshape(b, -1, n * d)
|
||||
x = self.o(x)
|
||||
x = self.dropout(x)
|
||||
return x
|
||||
|
||||
def flash_attn(self, x, context=None, mask=None, **kwargs):
|
||||
'''
|
||||
The implementation will be very slow when mask is not None,
|
||||
because we need rearange the x/context features according to mask.
|
||||
Args:
|
||||
x:
|
||||
context:
|
||||
mask:
|
||||
**kwargs:
|
||||
Returns: x
|
||||
'''
|
||||
context = x if context is None else context
|
||||
dtype = kwargs.get('dtype', torch.float16)
|
||||
q_lens = kwargs.get('q_lens', None)
|
||||
|
||||
# if mask is not None or q_lens is not None:
|
||||
# warnings.warn("Detected mask or q_lens is not None, "
|
||||
# "which will be very slow because of the x/context features' rearrangement,"
|
||||
# "please use FlashMultiHeadAttention instead.")
|
||||
def half(x):
|
||||
return x if x.dtype in [torch.float16, torch.bfloat16
|
||||
] else x.to(dtype)
|
||||
|
||||
b, n, d = x.size(0), self.num_heads, self.head_dim
|
||||
|
||||
q = self.q(x).view(b, -1, n, d) # [B, Lq, Nq, C1].
|
||||
k = self.k(context).view(b, -1, n, d) # [B, Lk, Nk, C1]
|
||||
v = self.v(context).view(
|
||||
b, -1, n, d) # [B, Lk, Nk, C2] Nq must be divisible by Nk.
|
||||
|
||||
assert q.device.type == 'cuda' and q.size(-1) <= 256
|
||||
lq, lk, out_dtype = int(q.size(1)), int(k.size(1)), q.dtype
|
||||
# preprocess query
|
||||
if q_lens is None:
|
||||
q_lens = torch.tensor([lq] * b,
|
||||
dtype=torch.int32).to(q.device,
|
||||
non_blocking=True)
|
||||
# q_lens = (q.flatten(2, ).bool() + 1).sum(dim=-1).bool().sum(dim=-1)
|
||||
q = half(q.flatten(0, 1))
|
||||
else:
|
||||
q = half(torch.cat([q_v[:q_l] for q_v, q_l in zip(q, q_lens)]))
|
||||
|
||||
# preprocess key, value
|
||||
if mask is None:
|
||||
k_lens = torch.tensor([lk] * b,
|
||||
dtype=torch.int32).to(k.device,
|
||||
non_blocking=True)
|
||||
# k_lens = (k.flatten(2, ).bool() + 1).sum(dim=-1).bool().sum(dim=-1)
|
||||
k = half(k.flatten(0, 1))
|
||||
v = half(v.flatten(0, 1))
|
||||
else:
|
||||
assert mask.ndim in [1, 2, 3]
|
||||
k_lens = mask if mask.ndim == 1 else mask.flatten(start_dim=1).sum(
|
||||
dim=-1)
|
||||
k = half(torch.cat([k_v[:k_l] for k_v, k_l in zip(k, k_lens)]))
|
||||
v = half(torch.cat([v_v[:v_l] for v_v, v_l in zip(v, k_lens)]))
|
||||
|
||||
x = flash_attn_varlen_func(
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
cu_seqlens_q=torch.cat([q_lens.new_zeros([1]),
|
||||
q_lens]).cumsum(0, dtype=torch.int32),
|
||||
cu_seqlens_k=torch.cat([k_lens.new_zeros([1]),
|
||||
k_lens]).cumsum(0, dtype=torch.int32),
|
||||
max_seqlen_q=int(torch.max(q_lens).cpu().numpy()),
|
||||
max_seqlen_k=int(torch.max(k_lens).cpu().numpy()),
|
||||
dropout_p=self.attn_drop.p,
|
||||
softmax_scale=self.softmax_scale,
|
||||
causal=self.causal,
|
||||
window_size=self.window_size, # -1 means infinite context window
|
||||
deterministic=self.deterministic).unflatten(0, (b, lq))
|
||||
x = x.type(out_dtype)
|
||||
x = x.flatten(2)
|
||||
# output
|
||||
x = self.o(x)
|
||||
x = self.dropout(x)
|
||||
return x
|
||||
|
||||
def pytorch_attn(self, x, context=None, mask=None, **kwargs):
|
||||
"""x: [B, L, C].
|
||||
context: [B, L', C'] or None.
|
||||
"""
|
||||
context = x if context is None else context
|
||||
b, n, d = x.size(0), self.num_heads, self.head_dim
|
||||
|
||||
# compute query, key, value
|
||||
q = self.q(x).view(b, -1, n, d)
|
||||
k = self.k(context).view(b, -1, n, d)
|
||||
v = self.v(context).view(b, -1, n, d)
|
||||
# attention bias
|
||||
attn_bias = x.new_zeros(b, n, q.size(1), k.size(1))
|
||||
if mask is not None:
|
||||
assert mask.ndim in [2, 3]
|
||||
mask = mask.view(b, 1, 1,
|
||||
-1) if mask.ndim == 2 else mask.unsqueeze(1)
|
||||
attn_bias = attn_bias.masked_fill_(mask == 0,
|
||||
torch.finfo(x.dtype).min).to(
|
||||
q.dtype)
|
||||
|
||||
# compute attention (T5 does not use scaling)
|
||||
attn = torch.einsum('binc,bjnc->bnij', q * self.scale,
|
||||
k * self.scale) + attn_bias
|
||||
attn = F.softmax(attn.float(), dim=-1).type_as(attn)
|
||||
x = torch.einsum('bnij,bjnc->binc', attn, v.float())
|
||||
# output
|
||||
x = x.reshape(b, -1, n * d)
|
||||
x = self.o(x)
|
||||
x = self.dropout(x)
|
||||
return x
|
||||
|
||||
def forward(self, x, context=None, mask=None, **kwargs):
|
||||
"""x: [B, L, C].
|
||||
context: [B, L', C'] or None.
|
||||
"""
|
||||
x = getattr(self, self.backend)(x,
|
||||
context=context,
|
||||
mask=mask,
|
||||
**kwargs)
|
||||
return x
|
||||
|
||||
|
||||
def flash_preprocess(x, context=None, q_mask=None, mask=None):
|
||||
context = x if context is None else context
|
||||
b, x_l, x_hidden_size = x.shape
|
||||
x = x.flatten(0, 1)
|
||||
if q_mask is None:
|
||||
q_lens = torch.tensor([x_l] * b,
|
||||
dtype=torch.int32).to(x.device,
|
||||
non_blocking=True)
|
||||
else:
|
||||
assert q_mask.ndim in [1, 2, 3]
|
||||
q_lens = q_mask if q_mask.ndim == 1 else q_mask.flatten(
|
||||
start_dim=1).sum(dim=-1)
|
||||
|
||||
mask_b, mask_l, mask_hidden_size = context.shape
|
||||
|
||||
if mask is None:
|
||||
mask_lens = torch.tensor([mask_l] * mask_b,
|
||||
dtype=torch.int32).to(context.device,
|
||||
non_blocking=True)
|
||||
else:
|
||||
assert mask.ndim in [1, 2, 3]
|
||||
mask_lens = mask if mask.ndim == 1 else mask.flatten(start_dim=1).sum(
|
||||
dim=-1)
|
||||
|
||||
return_data = {
|
||||
'x':
|
||||
x,
|
||||
'context':
|
||||
torch.cat([u[:v] for u, v in zip(context, mask_lens)]),
|
||||
'cu_seqlens_q':
|
||||
torch.cat([q_lens.new_zeros([1]), q_lens]).cumsum(0,
|
||||
dtype=torch.int32),
|
||||
'max_seqlen_q':
|
||||
int(torch.max(q_lens).cpu().numpy()),
|
||||
'cu_seqlens_k':
|
||||
torch.cat([mask_lens.new_zeros([1]),
|
||||
mask_lens]).cumsum(0, dtype=torch.int32),
|
||||
'max_seqlen_k':
|
||||
int(torch.max(mask_lens).cpu().numpy())
|
||||
}
|
||||
return return_data
|
||||
|
||||
|
||||
class FlashMultiHeadAttention(nn.Module):
|
||||
def __init__(self,
|
||||
dim,
|
||||
context_dim=None,
|
||||
num_heads=None,
|
||||
head_dim=None,
|
||||
attn_drop=0.0,
|
||||
qkv_bias=False,
|
||||
dropout=0.0,
|
||||
softmax_scale=None,
|
||||
causal=False,
|
||||
window_size=(-1, -1),
|
||||
deterministic=False,
|
||||
**block_kwargs):
|
||||
super().__init__()
|
||||
# consider head_dim first, then num_heads
|
||||
num_heads = dim // head_dim if head_dim else num_heads
|
||||
head_dim = dim // num_heads
|
||||
assert num_heads * head_dim == dim
|
||||
context_dim = context_dim or dim
|
||||
self.dim = dim
|
||||
self.context_dim = context_dim
|
||||
self.num_heads = num_heads
|
||||
self.head_dim = head_dim
|
||||
self.scale = math.pow(head_dim, -0.25)
|
||||
# layers
|
||||
self.q = nn.Linear(dim, dim, bias=qkv_bias)
|
||||
self.k = nn.Linear(context_dim, dim, bias=qkv_bias)
|
||||
self.v = nn.Linear(context_dim, dim, bias=qkv_bias)
|
||||
self.o = nn.Linear(dim, dim)
|
||||
self.dropout = nn.Dropout(dropout)
|
||||
self.attention_op = None
|
||||
self.attn_drop = nn.Dropout(attn_drop)
|
||||
self.softmax_scale = softmax_scale
|
||||
self.causal = causal
|
||||
self.window_size = window_size
|
||||
self.deterministic = deterministic
|
||||
|
||||
def forward(self,
|
||||
x,
|
||||
context=None,
|
||||
cu_seqlens_q=None,
|
||||
max_seqlen_q=None,
|
||||
cu_seqlens_k=None,
|
||||
max_seqlen_k=None,
|
||||
dtype=torch.float16,
|
||||
**kwargs):
|
||||
'''
|
||||
The implementation used the rearanaged x/context according to q_lens or k_lens.
|
||||
Args:
|
||||
x: [batch_size * max_seq_len or sum(q_lens) , heads, hidden_size].
|
||||
context: [batch_size * max_seq_len or sum(q_lens) , heads, hidden_size].
|
||||
cu_seqlens_q: cumsum of seq_q to index the postion of query in the batch.
|
||||
max_seqlen_q: max length of query.
|
||||
cu_seqlens_k: cumsum of seq_k to index the postion of key/value in the batch.
|
||||
max_seqlen_k: max length of key/value.
|
||||
dtype: the dtype for attention.
|
||||
**kwargs:
|
||||
Returns: x
|
||||
'''
|
||||
context = x if context is None else context
|
||||
|
||||
def half(x):
|
||||
return x if x.dtype in [torch.float16, torch.bfloat16
|
||||
] else x.to(dtype)
|
||||
|
||||
n, d, out_dtype = self.num_heads, self.head_dim, x.dtype
|
||||
q = self.q(x).view(-1, n, d) # [B * Lq, Nq, C1].
|
||||
k = self.k(context).view(-1, n, d) # [B * Lk, Nk, C1]
|
||||
v = self.v(context).view(
|
||||
-1, n, d) # [B * Lk, Nk, C2] Nq must be divisible by Nk.
|
||||
q, k, v = half(q), half(k), half(v)
|
||||
assert q.device.type == 'cuda' and d <= 256
|
||||
x = flash_attn_varlen_func(
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
cu_seqlens_q=cu_seqlens_q,
|
||||
cu_seqlens_k=cu_seqlens_k,
|
||||
max_seqlen_q=max_seqlen_q,
|
||||
max_seqlen_k=max_seqlen_k,
|
||||
dropout_p=self.attn_drop.p,
|
||||
softmax_scale=self.softmax_scale,
|
||||
causal=self.causal,
|
||||
window_size=self.window_size, # -1 means infinite context window
|
||||
deterministic=self.deterministic).unflatten(0, (x.shape[0], ))
|
||||
x = x.flatten(1).type(out_dtype)
|
||||
# output
|
||||
x = self.o(x)
|
||||
x = self.dropout(x)
|
||||
return x
|
||||
|
||||
|
||||
def multi_head_varlen_attention(q_img,
|
||||
k_img,
|
||||
v_img,
|
||||
q_txt,
|
||||
k_txt,
|
||||
v_txt,
|
||||
n,
|
||||
d,
|
||||
img_lens,
|
||||
txt_lens,
|
||||
dropout_p=0.0,
|
||||
flash_dtype=torch.bfloat16):
|
||||
'''
|
||||
q/k/v: b, s, n*d
|
||||
q_lens/k_lens: b,
|
||||
'''
|
||||
from flash_attn import flash_attn_varlen_func
|
||||
q_lens = k_lens = img_lens + txt_lens
|
||||
|
||||
cu_seqlens_q = torch.cat([q_lens.new_zeros([1]),
|
||||
q_lens]).cumsum(0, dtype=torch.int32)
|
||||
cu_seqlens_k = torch.cat([k_lens.new_zeros([1]),
|
||||
k_lens]).cumsum(0, dtype=torch.int32)
|
||||
max_seqlen_q = q_lens.max()
|
||||
max_seqlen_k = k_lens.max()
|
||||
|
||||
# concat img & txt for joint attention
|
||||
q = torch.cat([
|
||||
torch.cat([i[:i_len], t[:t_len]], dim=0)
|
||||
for i, i_len, t, t_len in zip(q_img, img_lens, q_txt, txt_lens)
|
||||
],
|
||||
dim=0).view(-1, n, d)
|
||||
|
||||
k = torch.cat([
|
||||
torch.cat([i[:i_len], t[:t_len]], dim=0)
|
||||
for i, i_len, t, t_len in zip(k_img, img_lens, k_txt, txt_lens)
|
||||
],
|
||||
dim=0).view(-1, n, d)
|
||||
|
||||
v = torch.cat([
|
||||
torch.cat([i[:i_len], t[:t_len]], dim=0)
|
||||
for i, i_len, t, t_len in zip(v_img, img_lens, v_txt, txt_lens)
|
||||
],
|
||||
dim=0).view(-1, n, d)
|
||||
|
||||
# attention
|
||||
dtype = q.dtype
|
||||
if dtype != flash_dtype:
|
||||
q = q.type(flash_dtype)
|
||||
k = k.type(flash_dtype)
|
||||
v = v.type(flash_dtype)
|
||||
|
||||
with amp.autocast():
|
||||
x = flash_attn_varlen_func(q=q,
|
||||
k=k,
|
||||
v=v,
|
||||
cu_seqlens_q=cu_seqlens_q,
|
||||
cu_seqlens_k=cu_seqlens_k,
|
||||
max_seqlen_q=max_seqlen_q,
|
||||
max_seqlen_k=max_seqlen_k,
|
||||
dropout_p=dropout_p).type(dtype)
|
||||
|
||||
return x, cu_seqlens_q
|
||||
|
||||
|
||||
class RMSNorm(nn.Module):
|
||||
def __init__(self, dim, eps=1e-6):
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.eps = eps
|
||||
self.weight = nn.Parameter(torch.ones(dim))
|
||||
|
||||
def forward(self, x):
|
||||
return self._norm(x.float()).type_as(x) * self.weight
|
||||
|
||||
def _norm(self, x):
|
||||
return x * torch.rsqrt(x.pow(2).mean(dim=-1, keepdim=True) + self.eps)
|
||||
|
||||
|
||||
class FullAttention(nn.Module):
|
||||
def __init__(self,
|
||||
dim,
|
||||
num_heads=None,
|
||||
head_dim=None,
|
||||
dropout=0.0,
|
||||
qkv_bias=False,
|
||||
qk_norm=False,
|
||||
eps=1e-6,
|
||||
flash_dtype=torch.bfloat16):
|
||||
|
||||
# consider head_dim first, then num_heads
|
||||
num_heads = dim // head_dim if head_dim else num_heads
|
||||
head_dim = dim // num_heads
|
||||
assert num_heads * head_dim == dim
|
||||
assert flash_dtype in (None, torch.float16, torch.bfloat16)
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.num_heads = num_heads
|
||||
self.head_dim = head_dim
|
||||
self.scale = math.pow(head_dim, -0.25)
|
||||
self.flash_dtype = flash_dtype
|
||||
# layers
|
||||
self.qkv_W = nn.Linear(dim, dim * 3, bias=qkv_bias)
|
||||
|
||||
self.out_proj = nn.Linear(dim, dim)
|
||||
self.dropout = nn.Dropout(dropout)
|
||||
|
||||
if qk_norm:
|
||||
from apex.normalization import FusedRMSNorm
|
||||
self.q_img_norm = FusedRMSNorm(head_dim, eps=eps)
|
||||
self.k_img_norm = FusedRMSNorm(head_dim, eps=eps)
|
||||
self.q_txt_norm = FusedRMSNorm(head_dim, eps=eps)
|
||||
self.k_txt_norm = FusedRMSNorm(head_dim, eps=eps)
|
||||
else:
|
||||
self.q_img_norm = nn.Identity()
|
||||
self.k_img_norm = nn.Identity()
|
||||
self.q_txt_norm = nn.Identity()
|
||||
self.k_txt_norm = nn.Identity()
|
||||
|
||||
def forward(self,
|
||||
img,
|
||||
txt,
|
||||
img_lens=None,
|
||||
txt_lens=None,
|
||||
padded_pos_index=None):
|
||||
'''
|
||||
img: B, L, C
|
||||
txt: B, L', C
|
||||
'''
|
||||
b, img_len, c = img.shape
|
||||
txt_len, n, d = txt.shape[1], self.num_heads, self.head_dim
|
||||
|
||||
# compute query, key, value
|
||||
img_txt = torch.cat([img, txt], dim=1)
|
||||
img_tokens, txt_tokens = torch.split(self.qkv_W(img_txt),
|
||||
[img_len, txt_len],
|
||||
dim=1)
|
||||
|
||||
q_img, k_img, v_img = img_tokens.chunk(3, dim=-1)
|
||||
q_txt, k_txt, v_txt = txt_tokens.chunk(3, dim=-1)
|
||||
|
||||
# multi-head qk norm
|
||||
q_img, k_img = q_img.view(b, -1, n, d), k_img.view(b, -1, n, d)
|
||||
q_txt, k_txt = q_txt.view(b, -1, n, d), k_txt.view(b, -1, n, d)
|
||||
q_img, q_txt = self.q_img_norm(q_img).view(
|
||||
b, -1, n * d), self.q_txt_norm(q_txt).view(b, -1, n * d)
|
||||
k_img, k_txt = self.k_img_norm(k_img).view(
|
||||
b, -1, n * d), self.k_txt_norm(k_txt).view(b, -1, n * d)
|
||||
|
||||
### add position
|
||||
q_img, k_img = apply_2d_rope(q_img, k_img, padded_pos_index, n, d)
|
||||
|
||||
# support varying length
|
||||
if img_lens is None:
|
||||
img_lens = torch.tensor([img.size(1)] * b,
|
||||
dtype=torch.int32,
|
||||
device=img.device)
|
||||
if txt_lens is None:
|
||||
txt_lens = torch.tensor([txt.size(1)] * b,
|
||||
dtype=torch.int32,
|
||||
device=txt.device)
|
||||
|
||||
# attention
|
||||
x, cu_seqlens_q = multi_head_varlen_attention(
|
||||
q_img,
|
||||
k_img,
|
||||
v_img,
|
||||
q_txt,
|
||||
k_txt,
|
||||
v_txt,
|
||||
n,
|
||||
d,
|
||||
img_lens,
|
||||
txt_lens,
|
||||
dropout_p=self.dropout.p if self.training else 0.0,
|
||||
flash_dtype=self.flash_dtype)
|
||||
|
||||
# output proj.
|
||||
x = x.reshape(-1, n * d)
|
||||
x = self.out_proj(x)
|
||||
x = self.dropout(x)
|
||||
|
||||
# split img & txt and padding to max_len
|
||||
img = pad_sequence(tuple([
|
||||
x[s:s + img_len] for s, e, img_len in zip(
|
||||
cu_seqlens_q[:-1], cu_seqlens_q[1:], img_lens)
|
||||
]),
|
||||
batch_first=True)
|
||||
txt = pad_sequence(tuple([
|
||||
x[s + img_len:e] for s, e, img_len in zip(
|
||||
cu_seqlens_q[:-1], cu_seqlens_q[1:], img_lens)
|
||||
]),
|
||||
batch_first=True)
|
||||
|
||||
return img, txt
|
||||
|
||||
|
||||
class FFNSwiGLU(nn.Module):
|
||||
def __init__(self, in_features, hidden_features):
|
||||
super().__init__()
|
||||
self.W1 = nn.Linear(in_features, hidden_features, bias=False)
|
||||
self.W2 = nn.Linear(in_features, hidden_features, bias=False)
|
||||
self.W3 = nn.Linear(hidden_features, in_features, bias=False)
|
||||
self.silu = nn.SiLU()
|
||||
|
||||
def forward(self, x):
|
||||
return self.W3(self.silu(self.W1(x)) * self.W2(x))
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
# Align results for different attention implementation
|
||||
torch.manual_seed(2023)
|
||||
hidden_dim = 4096
|
||||
q_weight = torch.randn((hidden_dim, hidden_dim))
|
||||
q_bias = torch.zeros((hidden_dim))
|
||||
k_weight = torch.randn((hidden_dim, hidden_dim))
|
||||
k_bias = torch.zeros((hidden_dim))
|
||||
v_weight = torch.randn((hidden_dim, hidden_dim))
|
||||
v_bias = torch.zeros((hidden_dim))
|
||||
o_weight = torch.randn((hidden_dim, hidden_dim))
|
||||
o_bias = torch.randn((hidden_dim))
|
||||
pytorch_attn = MultiHeadAttention(hidden_dim,
|
||||
context_dim=hidden_dim,
|
||||
num_heads=32,
|
||||
head_dim=None,
|
||||
attn_drop=0.0,
|
||||
dropout=0.0,
|
||||
backend='pytorch_attn')
|
||||
|
||||
pytorch_attn.load_state_dict({
|
||||
'q.weight': q_weight,
|
||||
'k.weight': k_weight,
|
||||
'v.weight': v_weight,
|
||||
'o.weight': o_weight,
|
||||
'o.bias': o_bias
|
||||
})
|
||||
pytorch_attn.to(0)
|
||||
|
||||
xformer_attn = MultiHeadAttention(hidden_dim,
|
||||
context_dim=hidden_dim,
|
||||
num_heads=32,
|
||||
head_dim=None,
|
||||
attn_drop=0.0,
|
||||
dropout=0.0,
|
||||
backend='xformer_attn')
|
||||
|
||||
xformer_attn.load_state_dict({
|
||||
'q.weight': q_weight,
|
||||
'k.weight': k_weight,
|
||||
'v.weight': v_weight,
|
||||
'o.weight': o_weight,
|
||||
'o.bias': o_bias
|
||||
})
|
||||
xformer_attn.to(0)
|
||||
|
||||
flash_attn = MultiHeadAttention(hidden_dim,
|
||||
context_dim=hidden_dim,
|
||||
num_heads=32,
|
||||
head_dim=None,
|
||||
attn_drop=0.0,
|
||||
dropout=0.0,
|
||||
backend='flash_attn',
|
||||
dtype=torch.float16)
|
||||
flash_attn.load_state_dict({
|
||||
'q.weight': q_weight,
|
||||
'k.weight': k_weight,
|
||||
'v.weight': v_weight,
|
||||
'o.weight': o_weight,
|
||||
'o.bias': o_bias
|
||||
})
|
||||
flash_attn.to(0)
|
||||
|
||||
improved_flash_attn = FlashMultiHeadAttention(hidden_dim,
|
||||
context_dim=hidden_dim,
|
||||
num_heads=32,
|
||||
head_dim=None,
|
||||
attn_drop=0.0,
|
||||
dropout=0.0,
|
||||
backend='flash_attn',
|
||||
dtype=torch.float16)
|
||||
|
||||
improved_flash_attn.load_state_dict({
|
||||
'q.weight': q_weight,
|
||||
'k.weight': k_weight,
|
||||
'v.weight': v_weight,
|
||||
'o.weight': o_weight,
|
||||
'o.bias': o_bias
|
||||
})
|
||||
improved_flash_attn.to(0)
|
||||
|
||||
batch_size = 1
|
||||
query_length = 1024
|
||||
key_length = 1024
|
||||
# mask = None
|
||||
run_num = 10
|
||||
torch.cuda.empty_cache()
|
||||
x = torch.randn((batch_size, query_length, hidden_dim)).to(0)
|
||||
context = torch.randn((batch_size, key_length, hidden_dim)).to(0)
|
||||
# mask = torch.cat([torch.ones((batch_size, 80)), torch.zeros((batch_size, key_length - 80))], dim=1).long().to(0)
|
||||
# mask = torch.randint(1, key_length, [batch_size]).to(0)
|
||||
mask = None
|
||||
st = time.time()
|
||||
for i in tqdm(range(run_num)):
|
||||
pytorch_res = pytorch_attn(x.clone(), context.clone(),
|
||||
mask.clone() if mask is not None else mask)
|
||||
if i == run_num - 1:
|
||||
free_mem, total_mem = torch.cuda.mem_get_info(0)
|
||||
free_mem, total_mem = free_mem / (1024**3), total_mem / (1024**3)
|
||||
mem_msg = f'GPU {0}: free mem {free_mem:.3f}G, total mem {total_mem:.3f}G \n'
|
||||
pytorch_res_data = pytorch_res.clone().detach().cpu()
|
||||
print('pytorch attn ', mem_msg,
|
||||
f'Cost time per time {(time.time() - st) / run_num}s')
|
||||
#
|
||||
torch.cuda.empty_cache()
|
||||
st = time.time()
|
||||
for i in tqdm(range(run_num)):
|
||||
xformer_res = xformer_attn(x.clone(), context.clone(),
|
||||
mask.clone() if mask is not None else mask)
|
||||
if i == run_num - 1:
|
||||
free_mem, total_mem = torch.cuda.mem_get_info(0)
|
||||
free_mem, total_mem = free_mem / (1024**3), total_mem / (1024**3)
|
||||
mem_msg = f'GPU {0}: free mem {free_mem:.3f}G, total mem {total_mem:.3f}G \n'
|
||||
xformer_res_data = xformer_res.clone().detach().cpu()
|
||||
print('xformer attn ', mem_msg,
|
||||
f'Cost time per time {(time.time() - st) / run_num}s')
|
||||
#
|
||||
torch.cuda.empty_cache()
|
||||
# mask = None
|
||||
st = time.time()
|
||||
for i in tqdm(range(run_num)):
|
||||
flash_res = flash_attn(x.clone(), context.clone(),
|
||||
mask.clone() if mask is not None else mask)
|
||||
if i == run_num - 1:
|
||||
free_mem, total_mem = torch.cuda.mem_get_info(0)
|
||||
free_mem, total_mem = free_mem / (1024**3), total_mem / (1024**3)
|
||||
mem_msg = f'GPU {0}: free mem {free_mem:.3f}G, total mem {total_mem:.3f}G \n'
|
||||
flash_res_data = flash_res.clone().detach().cpu()
|
||||
print('flash attn ', mem_msg,
|
||||
f'Cost time per time {(time.time() - st) / run_num}s')
|
||||
|
||||
# recommend this style for multi blocks to save the preprocess time.
|
||||
|
||||
flash_input = flash_preprocess(
|
||||
x.clone(),
|
||||
context.clone(),
|
||||
mask=mask.clone() if mask is not None else mask)
|
||||
st = time.time()
|
||||
for i in tqdm(range(run_num)):
|
||||
improved_flash_res_v1 = improved_flash_attn(**flash_input).reshape(
|
||||
(batch_size, -1, hidden_dim))
|
||||
if i == 0:
|
||||
improved_flash_res_v1_data = improved_flash_res_v1.clone().detach(
|
||||
).cpu()
|
||||
if i == run_num - 1:
|
||||
free_mem, total_mem = torch.cuda.mem_get_info(0)
|
||||
free_mem, total_mem = free_mem / (1024**3), total_mem / (1024**3)
|
||||
mem_msg = f'GPU {0}: free mem {free_mem:.3f}G, total mem {total_mem:.3f}G \n'
|
||||
print('improved flash attn ', mem_msg,
|
||||
f'Cost time per time {(time.time() - st) / run_num}s')
|
||||
#
|
||||
print(pytorch_res_data, xformer_res_data, flash_res_data,
|
||||
improved_flash_res_v1_data)
|
||||
print(pytorch_res_data.shape, xformer_res_data.shape, flash_res_data.shape,
|
||||
improved_flash_res_v1_data.shape)
|
||||
print(
|
||||
torch.sum(pytorch_res_data) / (batch_size * query_length * hidden_dim),
|
||||
torch.sum(xformer_res_data) / (batch_size * query_length * hidden_dim),
|
||||
torch.sum(flash_res_data) / (batch_size * query_length * hidden_dim),
|
||||
torch.sum(improved_flash_res_v1_data) /
|
||||
(batch_size * query_length * hidden_dim))
|
||||
@@ -0,0 +1,303 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
# All rights reserved.
|
||||
# This file contains code that is adapted from
|
||||
# timm: https://github.com/huggingface/pytorch-image-models
|
||||
# pixart: https://github.com/PixArt-alpha/PixArt-alpha
|
||||
import math
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from einops import rearrange
|
||||
|
||||
from scepter.modules.model.backbone.transformer.attention import drop_path
|
||||
|
||||
|
||||
def modulate(x, shift, scale, unsqueeze=False):
|
||||
if unsqueeze:
|
||||
return x * (1 + scale.unsqueeze(1)) + shift.unsqueeze(1)
|
||||
else:
|
||||
return x * (1 + scale) + shift
|
||||
|
||||
|
||||
class DropPath(nn.Module):
|
||||
"""Drop paths (Stochastic Depth) per sample (when applied in main path of residual blocks).
|
||||
"""
|
||||
def __init__(self, drop_prob=None):
|
||||
super(DropPath, self).__init__()
|
||||
self.drop_prob = drop_prob
|
||||
|
||||
def forward(self, x):
|
||||
return drop_path(x, self.drop_prob, self.training)
|
||||
|
||||
|
||||
class MaskFinalLayer(nn.Module):
|
||||
"""
|
||||
The final layer of PixArt.
|
||||
"""
|
||||
def __init__(self, final_hidden_size, c_emb_size, patch_size,
|
||||
out_channels):
|
||||
super().__init__()
|
||||
self.norm_final = nn.LayerNorm(final_hidden_size,
|
||||
elementwise_affine=False,
|
||||
eps=1e-6)
|
||||
self.linear = nn.Linear(final_hidden_size,
|
||||
patch_size * patch_size * out_channels,
|
||||
bias=True)
|
||||
self.adaLN_modulation = nn.Sequential(
|
||||
nn.SiLU(), nn.Linear(c_emb_size, 2 * final_hidden_size, bias=True))
|
||||
|
||||
def forward(self, x, t):
|
||||
shift, scale = self.adaLN_modulation(t).chunk(2, dim=1)
|
||||
x = modulate(self.norm_final(x), shift, scale, unsqueeze=True)
|
||||
x = self.linear(x)
|
||||
return x
|
||||
|
||||
|
||||
class DecoderLayer(nn.Module):
|
||||
"""
|
||||
The final layer of PixArt.
|
||||
"""
|
||||
def __init__(self, hidden_size, decoder_hidden_size):
|
||||
super().__init__()
|
||||
self.norm_decoder = nn.LayerNorm(hidden_size,
|
||||
elementwise_affine=False,
|
||||
eps=1e-6)
|
||||
self.linear = nn.Linear(hidden_size, decoder_hidden_size, bias=True)
|
||||
self.adaLN_modulation = nn.Sequential(
|
||||
nn.SiLU(), nn.Linear(hidden_size, 2 * hidden_size, bias=True))
|
||||
|
||||
def forward(self, x, t):
|
||||
shift, scale = self.adaLN_modulation(t).chunk(2, dim=1)
|
||||
x = modulate(self.norm_decoder(x), shift, scale, unsqueeze=True)
|
||||
x = self.linear(x)
|
||||
return x
|
||||
|
||||
|
||||
class TimestepEmbedder(nn.Module):
|
||||
"""
|
||||
Embeds scalar timesteps into vector representations.
|
||||
"""
|
||||
def __init__(self, hidden_size, frequency_embedding_size=256):
|
||||
super().__init__()
|
||||
self.mlp = nn.Sequential(
|
||||
nn.Linear(frequency_embedding_size, hidden_size, bias=True),
|
||||
nn.SiLU(),
|
||||
nn.Linear(hidden_size, hidden_size, bias=True),
|
||||
)
|
||||
self.frequency_embedding_size = frequency_embedding_size
|
||||
|
||||
@staticmethod
|
||||
def timestep_embedding(t, dim, max_period=10000):
|
||||
"""
|
||||
Create sinusoidal timestep embeddings.
|
||||
:param t: a 1-D Tensor of N indices, one per batch element.
|
||||
These may be fractional.
|
||||
:param dim: the dimension of the output.
|
||||
:param max_period: controls the minimum frequency of the embeddings.
|
||||
:return: an (N, D) Tensor of positional embeddings.
|
||||
"""
|
||||
# https://github.com/openai/glide-text2im/blob/main/glide_text2im/nn.py
|
||||
half = dim // 2
|
||||
freqs = torch.exp(
|
||||
-math.log(max_period) *
|
||||
torch.arange(start=0, end=half, dtype=torch.float32) /
|
||||
half).to(device=t.device)
|
||||
args = t[:, None].float() * freqs[None]
|
||||
embedding = torch.cat([torch.cos(args), torch.sin(args)], dim=-1)
|
||||
if dim % 2:
|
||||
embedding = torch.cat(
|
||||
[embedding, torch.zeros_like(embedding[:, :1])], dim=-1)
|
||||
return embedding
|
||||
|
||||
def forward(self, t):
|
||||
t_freq = self.timestep_embedding(t, self.frequency_embedding_size)
|
||||
t_emb = self.mlp(t_freq)
|
||||
return t_emb
|
||||
|
||||
|
||||
class SizeEmbedder(TimestepEmbedder):
|
||||
"""
|
||||
Embeds scalar timesteps into vector representations.
|
||||
"""
|
||||
def __init__(self, hidden_size, frequency_embedding_size=256):
|
||||
super().__init__(hidden_size=hidden_size,
|
||||
frequency_embedding_size=frequency_embedding_size)
|
||||
self.mlp = nn.Sequential(
|
||||
nn.Linear(frequency_embedding_size, hidden_size, bias=True),
|
||||
nn.SiLU(),
|
||||
nn.Linear(hidden_size, hidden_size, bias=True),
|
||||
)
|
||||
self.frequency_embedding_size = frequency_embedding_size
|
||||
self.outdim = hidden_size
|
||||
|
||||
def forward(self, s, bs):
|
||||
if s.ndim == 1:
|
||||
s = s[:, None]
|
||||
assert s.ndim == 2
|
||||
if s.shape[0] != bs:
|
||||
s = s.repeat(bs // s.shape[0], 1)
|
||||
assert s.shape[0] == bs
|
||||
b, dims = s.shape[0], s.shape[1]
|
||||
s = rearrange(s, 'b d -> (b d)')
|
||||
s_freq = self.timestep_embedding(s, self.frequency_embedding_size).to(
|
||||
self.dtype)
|
||||
s_emb = self.mlp(s_freq)
|
||||
s_emb = rearrange(s_emb,
|
||||
'(b d) d2 -> b (d d2)',
|
||||
b=b,
|
||||
d=dims,
|
||||
d2=self.outdim)
|
||||
return s_emb
|
||||
|
||||
@property
|
||||
def dtype(self):
|
||||
# 返回模型参数的数据类型
|
||||
return next(self.parameters()).dtype
|
||||
|
||||
|
||||
class LabelEmbedder(nn.Module):
|
||||
"""
|
||||
Embeds class labels into vector representations. Also handles label dropout for classifier-free guidance.
|
||||
"""
|
||||
def __init__(self, num_classes, hidden_size, dropout_prob):
|
||||
super().__init__()
|
||||
use_cfg_embedding = dropout_prob > 0
|
||||
self.embedding_table = nn.Embedding(num_classes + use_cfg_embedding,
|
||||
hidden_size)
|
||||
self.num_classes = num_classes
|
||||
self.dropout_prob = dropout_prob
|
||||
|
||||
def token_drop(self, labels, force_drop_ids=None):
|
||||
"""
|
||||
Drops labels to enable classifier-free guidance.
|
||||
"""
|
||||
if force_drop_ids is None:
|
||||
drop_ids = torch.rand(labels.shape[0]).cuda() < self.dropout_prob
|
||||
else:
|
||||
drop_ids = force_drop_ids == 1
|
||||
labels = torch.where(drop_ids, self.num_classes, labels)
|
||||
return labels
|
||||
|
||||
def forward(self, labels, train, force_drop_ids=None):
|
||||
use_dropout = self.dropout_prob > 0
|
||||
if (train and use_dropout) or (force_drop_ids is not None):
|
||||
labels = self.token_drop(labels, force_drop_ids)
|
||||
return self.embedding_table(labels)
|
||||
|
||||
|
||||
class CaptionEmbedder(nn.Module):
|
||||
"""
|
||||
Embeds class labels into vector representations. Also handles label dropout for classifier-free guidance.
|
||||
"""
|
||||
def __init__(self,
|
||||
in_channels,
|
||||
hidden_size,
|
||||
uncond_prob,
|
||||
act_layer=nn.GELU(approximate='tanh'),
|
||||
token_num=120):
|
||||
super().__init__()
|
||||
self.y_proj = Mlp(in_features=in_channels,
|
||||
hidden_features=hidden_size,
|
||||
out_features=hidden_size,
|
||||
act_layer=act_layer,
|
||||
drop=0)
|
||||
self.register_buffer(
|
||||
'y_embedding',
|
||||
nn.Parameter(
|
||||
torch.randn(token_num, in_channels) / in_channels**0.5))
|
||||
self.uncond_prob = uncond_prob
|
||||
|
||||
def token_drop(self, caption, force_drop_ids=None):
|
||||
"""
|
||||
Drops labels to enable classifier-free guidance.
|
||||
"""
|
||||
if force_drop_ids is None:
|
||||
drop_ids = torch.rand(caption.shape[0]).cuda() < self.uncond_prob
|
||||
else:
|
||||
drop_ids = force_drop_ids == 1
|
||||
caption = torch.where(drop_ids[:, None, None], self.y_embedding,
|
||||
caption)
|
||||
return caption
|
||||
|
||||
def forward(self, caption, train, force_drop_ids=None):
|
||||
if train:
|
||||
assert caption.shape[1:] == self.y_embedding.shape
|
||||
use_dropout = self.uncond_prob > 0
|
||||
if (train and use_dropout) or (force_drop_ids is not None):
|
||||
caption = self.token_drop(caption, force_drop_ids)
|
||||
caption = self.y_proj(caption)
|
||||
return caption
|
||||
|
||||
|
||||
class CaptionEmbedderDoubleBr(nn.Module):
|
||||
"""
|
||||
Embeds class labels into vector representations. Also handles label dropout for classifier-free guidance.
|
||||
"""
|
||||
def __init__(self,
|
||||
in_channels,
|
||||
hidden_size,
|
||||
uncond_prob,
|
||||
act_layer=nn.GELU(approximate='tanh'),
|
||||
token_num=120):
|
||||
super().__init__()
|
||||
self.proj = Mlp(in_features=in_channels,
|
||||
hidden_features=hidden_size,
|
||||
out_features=hidden_size,
|
||||
act_layer=act_layer,
|
||||
drop=0)
|
||||
self.embedding = nn.Parameter(torch.randn(1, in_channels) / 10**0.5)
|
||||
self.y_embedding = nn.Parameter(
|
||||
torch.randn(token_num, in_channels) / 10**0.5)
|
||||
self.uncond_prob = uncond_prob
|
||||
|
||||
def token_drop(self, global_caption, caption, force_drop_ids=None):
|
||||
"""
|
||||
Drops labels to enable classifier-free guidance.
|
||||
"""
|
||||
if force_drop_ids is None:
|
||||
drop_ids = torch.rand(
|
||||
global_caption.shape[0]).cuda() < self.uncond_prob
|
||||
else:
|
||||
drop_ids = force_drop_ids == 1
|
||||
global_caption = torch.where(drop_ids[:, None], self.embedding,
|
||||
global_caption)
|
||||
caption = torch.where(drop_ids[:, None, None, None], self.y_embedding,
|
||||
caption)
|
||||
return global_caption, caption
|
||||
|
||||
def forward(self, caption, train, force_drop_ids=None):
|
||||
assert caption.shape[2:] == self.y_embedding.shape
|
||||
global_caption = caption.mean(dim=2).squeeze()
|
||||
use_dropout = self.uncond_prob > 0
|
||||
if (train and use_dropout) or (force_drop_ids is not None):
|
||||
global_caption, caption = self.token_drop(global_caption, caption,
|
||||
force_drop_ids)
|
||||
y_embed = self.proj(global_caption)
|
||||
return y_embed, caption
|
||||
|
||||
|
||||
class Mlp(nn.Module):
|
||||
""" MLP as used in Vision Transformer, MLP-Mixer and related networks
|
||||
"""
|
||||
def __init__(self,
|
||||
in_features,
|
||||
hidden_features=None,
|
||||
out_features=None,
|
||||
act_layer=nn.GELU,
|
||||
drop=0.):
|
||||
super().__init__()
|
||||
out_features = out_features or in_features
|
||||
hidden_features = hidden_features or in_features
|
||||
self.fc1 = nn.Linear(in_features, hidden_features)
|
||||
self.act = act_layer()
|
||||
self.fc2 = nn.Linear(hidden_features, out_features)
|
||||
self.drop = nn.Dropout(drop)
|
||||
|
||||
def forward(self, x):
|
||||
x = self.fc1(x)
|
||||
x = self.act(x)
|
||||
x = self.drop(x)
|
||||
x = self.fc2(x)
|
||||
x = self.drop(x)
|
||||
return x
|
||||
@@ -0,0 +1,54 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
# All rights reserved.
|
||||
# This file contains code that is adapted from
|
||||
# timm: https://github.com/huggingface/pytorch-image-models
|
||||
# pixart: https://github.com/PixArt-alpha/PixArt-alpha
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
|
||||
class PatchEmbed(nn.Module):
|
||||
""" 2D Image to Patch Embedding
|
||||
"""
|
||||
def __init__(
|
||||
self,
|
||||
patch_size=16,
|
||||
in_chans=3,
|
||||
embed_dim=768,
|
||||
norm_layer=None,
|
||||
flatten=True,
|
||||
bias=True,
|
||||
):
|
||||
super().__init__()
|
||||
self.flatten = flatten
|
||||
self.proj = nn.Conv2d(in_chans,
|
||||
embed_dim,
|
||||
kernel_size=patch_size,
|
||||
stride=patch_size,
|
||||
bias=bias)
|
||||
self.norm = norm_layer(embed_dim) if norm_layer else nn.Identity()
|
||||
|
||||
def forward(self, x):
|
||||
x = self.proj(x)
|
||||
if self.flatten:
|
||||
x = x.flatten(2).transpose(1, 2) # BCHW -> BNC
|
||||
x = self.norm(x)
|
||||
return x
|
||||
|
||||
|
||||
def unpatchify(x, h, w, c, p_h, p_w):
|
||||
'''
|
||||
Args:
|
||||
x: input tensor for unpatchified with shape as (N, T, patch_size**2 * C).
|
||||
h: tokens' number align height
|
||||
w: tokens' number align width
|
||||
c: output channels
|
||||
p_h: patch size for h
|
||||
p_w: patch size for w
|
||||
Returns: unpatchified imgs with shape as (N, H, W, C)
|
||||
'''
|
||||
assert h * w == x.shape[1]
|
||||
x = x.reshape(shape=(x.shape[0], h, w, p_h, p_w, c))
|
||||
x = torch.einsum('nhwpqc->nchpwq', x)
|
||||
return x.reshape(shape=(x.shape[0], c, h * p_h, w * p_w))
|
||||
@@ -0,0 +1,129 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
# All rights reserved.
|
||||
# This file contains code that is adapted from
|
||||
# timm: https://github.com/huggingface/pytorch-image-models
|
||||
# pixart: https://github.com/PixArt-alpha/PixArt-alpha
|
||||
from itertools import repeat as iter_repeat
|
||||
from typing import Iterable
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
|
||||
def _ntuple(n):
|
||||
def parse(x):
|
||||
if isinstance(x, Iterable) and not isinstance(x, str):
|
||||
return x
|
||||
return tuple(iter_repeat(x, n))
|
||||
|
||||
return parse
|
||||
|
||||
|
||||
to_1tuple = _ntuple(1)
|
||||
to_2tuple = _ntuple(2)
|
||||
|
||||
|
||||
def get_2d_sincos_pos_embed(embed_dim,
|
||||
grid_size,
|
||||
cls_token=False,
|
||||
extra_tokens=0,
|
||||
lewei_scale=1.0,
|
||||
base_h_size=16.,
|
||||
base_w_size=16):
|
||||
"""
|
||||
grid_size: int of the grid height and width
|
||||
return:
|
||||
pos_embed: [grid_size*grid_size, embed_dim] or [1+grid_size*grid_size, embed_dim] (w/ or w/o cls_token)
|
||||
"""
|
||||
if isinstance(grid_size, int):
|
||||
grid_size = to_2tuple(grid_size)
|
||||
grid_h = np.arange(grid_size[0], dtype=np.float32) / (
|
||||
grid_size[0] / base_h_size) / lewei_scale
|
||||
grid_w = np.arange(grid_size[1], dtype=np.float32) / (
|
||||
grid_size[1] / base_w_size) / lewei_scale
|
||||
grid = np.meshgrid(grid_w, grid_h) # here w goes first
|
||||
grid = np.stack(grid, axis=0)
|
||||
grid = grid.reshape([2, 1, grid_size[1], grid_size[0]])
|
||||
|
||||
pos_embed = get_2d_sincos_pos_embed_from_grid(embed_dim, grid)
|
||||
if cls_token and extra_tokens > 0:
|
||||
pos_embed = np.concatenate(
|
||||
[np.zeros([extra_tokens, embed_dim]), pos_embed], axis=0)
|
||||
return pos_embed
|
||||
|
||||
|
||||
def get_2d_sincos_pos_embed_from_grid(embed_dim, grid):
|
||||
assert embed_dim % 2 == 0
|
||||
|
||||
# use half of dimensions to encode grid_h
|
||||
emb_h = get_1d_sincos_pos_embed_from_grid(embed_dim // 2,
|
||||
grid[0]) # (H*W, D/2)
|
||||
emb_w = get_1d_sincos_pos_embed_from_grid(embed_dim // 2,
|
||||
grid[1]) # (H*W, D/2)
|
||||
|
||||
return np.concatenate([emb_h, emb_w], axis=1)
|
||||
|
||||
|
||||
def get_1d_sincos_pos_embed_from_grid(embed_dim, pos):
|
||||
"""
|
||||
embed_dim: output dimension for each position
|
||||
pos: a list of positions to be encoded: size (M,)
|
||||
out: (M, D)
|
||||
"""
|
||||
assert embed_dim % 2 == 0
|
||||
omega = np.arange(embed_dim // 2, dtype=np.float64)
|
||||
omega /= embed_dim / 2.
|
||||
omega = 1. / 10000**omega # (D/2,)
|
||||
|
||||
pos = pos.reshape(-1) # (M,)
|
||||
out = np.einsum('m,d->md', pos, omega) # (M, D/2), outer product
|
||||
|
||||
emb_sin = np.sin(out) # (M, D/2)
|
||||
emb_cos = np.cos(out) # (M, D/2)
|
||||
return np.concatenate([emb_sin, emb_cos], axis=1)
|
||||
|
||||
|
||||
def apply_2d_rope(xq,
|
||||
xk,
|
||||
padded_pos_index,
|
||||
num_head,
|
||||
head_dim,
|
||||
rotary_base=10000):
|
||||
'''
|
||||
x query/key: [b, seq, num_head*head_dim]
|
||||
padded_pos_index: [b, seq, 2]
|
||||
'''
|
||||
b = xq.shape[0]
|
||||
assert head_dim % 4 == 0, 'the 2d_rope dims should be divided by 4'
|
||||
rope_dim = head_dim // 2 # 2d_rope_dim, 1d_rope_dim = head_dim
|
||||
# 1. theta_d = b ** (-2d/D)
|
||||
theta = 1.0 / (rotary_base**(
|
||||
torch.arange(0, rope_dim, 2)[:(rope_dim // 2)].float() / rope_dim))
|
||||
# 2. [h * Theta || w * Theta]
|
||||
theta = theta.to(xq.device).expand(b, 1, rope_dim // 2)
|
||||
freqs_h = torch.bmm(padded_pos_index[:, :, :1],
|
||||
theta).float() # h * \theta
|
||||
freqs_w = torch.bmm(padded_pos_index[:, :, 1:],
|
||||
theta).float() # w * \theta
|
||||
freqs = torch.cat([freqs_h, freqs_w], dim=2).repeat(1, 1,
|
||||
num_head) # multi-head
|
||||
# 3. as_complex for complex multiply
|
||||
# if freqs = [x, y] then freqs_cis = [cos(x) + sin(x)i, cos(y) + sin(y)i]
|
||||
freqs_cis = torch.polar(
|
||||
torch.ones_like(freqs),
|
||||
freqs) # torch.polar(abs, angle)=> abs⋅cos(angle)+abs⋅sin(angle)⋅j
|
||||
# xq.shape = [b, seq_len, dim]
|
||||
# xq_.shape = [b, seq_len, dim // 2, 2]
|
||||
xq_ = xq.float().reshape(*xq.shape[:-1], -1, 2)
|
||||
xk_ = xk.float().reshape(*xk.shape[:-1], -1, 2)
|
||||
# 转为复数域
|
||||
xq_ = torch.view_as_complex(
|
||||
xq_) # [b, seq_len, dim // 2, 2]=>xq.shape = [b, seq_len, dim]
|
||||
xk_ = torch.view_as_complex(xk_)
|
||||
# 4. complex multiply and as real
|
||||
# xq_out.shape = [b, seq_len, dim]
|
||||
xq_out = torch.view_as_real(xq_ * freqs_cis).flatten(
|
||||
2) # point_wise mul, then flatten eg[[1,2],[3,4],[5,6]]->[1,2,3,4,5,6]
|
||||
xk_out = torch.view_as_real(xk_ * freqs_cis).flatten(2)
|
||||
return xq_out.type_as(xq), xk_out.type_as(xk)
|
||||
@@ -2,6 +2,6 @@
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
|
||||
from scepter.modules.model.embedder.embedder import (
|
||||
ConcatTimestepEmbedderND, FrozenCLIPEmbedder, FrozenOpenCLIPEmbedder,
|
||||
FrozenOpenCLIPEmbedder2, GeneralConditioner, IPAdapterPlusEmbedder,
|
||||
RefCrossEmbedder)
|
||||
ConcatTimestepEmbedderND, FrozenCLIPEmbedder, FrozenCLIPEmbedder2,
|
||||
FrozenOpenCLIPEmbedder, FrozenOpenCLIPEmbedder2, GeneralConditioner,
|
||||
IPAdapterPlusEmbedder, RefCrossEmbedder, SD3TextEmbedder, T5EmbedderHF)
|
||||
|
||||
@@ -11,20 +11,22 @@ import torch
|
||||
import torch.nn as nn
|
||||
import torch.utils.dlpack
|
||||
from einops import rearrange
|
||||
# to check
|
||||
from scepter.modules.model.backbone.unet.unet_utils import Timestep
|
||||
from scepter.modules.model.embedder.base_embedder import BaseEmbedder
|
||||
from scepter.modules.model.embedder.resampler import Resampler
|
||||
from scepter.modules.model.registry import EMBEDDERS
|
||||
from scepter.modules.model.tokenizer.tokenizer_component import (
|
||||
basic_clean, canonicalize, heavy_clean, whitespace_clean)
|
||||
from scepter.modules.model.utils.basic_utils import expand_dims_like
|
||||
from scepter.modules.utils.config import dict_to_yaml
|
||||
from scepter.modules.utils.distribute import we
|
||||
from scepter.modules.utils.file_system import FS
|
||||
from torch.utils.checkpoint import checkpoint
|
||||
|
||||
from .base_embedder import BaseEmbedder
|
||||
from .resampler import Resampler
|
||||
|
||||
try:
|
||||
from transformers import CLIPTextModel, CLIPTokenizer, CLIPVisionModelWithProjection
|
||||
from transformers import (CLIPTextModel, CLIPTokenizer,
|
||||
CLIPVisionModelWithProjection, AutoTokenizer,
|
||||
T5EncoderModel, CLIPTextModelWithProjection)
|
||||
except Exception as e:
|
||||
warnings.warn(
|
||||
f'Import transformers error, please deal with this problem: {e}')
|
||||
@@ -117,7 +119,6 @@ class FrozenCLIPEmbedder(BaseEmbedder):
|
||||
for param in self.parameters():
|
||||
param.requires_grad = False
|
||||
|
||||
# @torch.no_grad()
|
||||
def _forward(self, text):
|
||||
batch_encoding = self.tokenizer(text,
|
||||
truncation=True,
|
||||
@@ -779,3 +780,296 @@ class GeneralConditioner(BaseEmbedder):
|
||||
__class__.__name__,
|
||||
GeneralConditioner.para_dict,
|
||||
set_name=True)
|
||||
|
||||
|
||||
@EMBEDDERS.register_class()
|
||||
class T5EmbedderHF(BaseEmbedder):
|
||||
"""
|
||||
Uses the OpenCLIP transformer encoder for text
|
||||
"""
|
||||
"""
|
||||
Uses the OpenCLIP transformer encoder for text
|
||||
"""
|
||||
para_dict = {
|
||||
'PRETRAINED_MODEL': {
|
||||
'value':
|
||||
'google/umt5-small',
|
||||
'description':
|
||||
'Pretrained Model for umt5, modelcard path or local path.'
|
||||
},
|
||||
'TOKENIZER_PATH': {
|
||||
'value': 'google/umt5-small',
|
||||
'description':
|
||||
'Tokenizer Path for umt5, modelcard path or local path.'
|
||||
},
|
||||
'FREEZE': {
|
||||
'value': True,
|
||||
'description': ''
|
||||
},
|
||||
'USE_GRAD': {
|
||||
'value': False,
|
||||
'description': 'Compute grad or not.'
|
||||
},
|
||||
'CLEAN': {
|
||||
'value':
|
||||
'whitespace',
|
||||
'description':
|
||||
'Set the clean strtegy for tokenizer, used when TOKENIZER_PATH is not None.'
|
||||
},
|
||||
'LAYER': {
|
||||
'value': 'last',
|
||||
'description': ''
|
||||
},
|
||||
'LEGACY': {
|
||||
'value':
|
||||
True,
|
||||
'description':
|
||||
'Whether use legacy returnd feature or not ,default True.'
|
||||
}
|
||||
}
|
||||
|
||||
def __init__(self, cfg, logger=None):
|
||||
super().__init__(cfg, logger=logger)
|
||||
pretrained_path = cfg.get('PRETRAINED_MODEL', None)
|
||||
t5_dtype = cfg.get('T5_DTYPE', None)
|
||||
assert pretrained_path
|
||||
with FS.get_dir_to_local_dir(pretrained_path,
|
||||
wait_finish=True) as local_path:
|
||||
if t5_dtype is not None:
|
||||
self.model = T5EncoderModel.from_pretrained(
|
||||
local_path, torch_dtype=getattr(torch, t5_dtype))
|
||||
else:
|
||||
self.model = T5EncoderModel.from_pretrained(local_path)
|
||||
tokenizer_path = cfg.get('TOKENIZER_PATH', None)
|
||||
self.length = cfg.get('LENGTH', 77)
|
||||
if tokenizer_path:
|
||||
self.tokenize_kargs = {'return_tensors': 'pt'}
|
||||
with FS.get_dir_to_local_dir(tokenizer_path,
|
||||
wait_finish=True) as local_path:
|
||||
self.tokenizer = AutoTokenizer.from_pretrained(local_path)
|
||||
if self.length is not None:
|
||||
self.tokenize_kargs.update({
|
||||
'padding': 'max_length',
|
||||
'truncation': True,
|
||||
'max_length': self.length
|
||||
})
|
||||
self.eos_token = self.tokenizer(
|
||||
self.tokenizer.eos_token)['input_ids'][0]
|
||||
else:
|
||||
self.tokenizer = None
|
||||
self.tokenize_kargs = {}
|
||||
|
||||
self.use_grad = cfg.get('USE_GRAD', False)
|
||||
self.clean = cfg.get('CLEAN', 'whitespace')
|
||||
|
||||
def freeze(self):
|
||||
self.model = self.model.eval()
|
||||
for param in self.parameters():
|
||||
param.requires_grad = False
|
||||
|
||||
# encode && encode_text
|
||||
def forward(self, tokens, return_mask=False):
|
||||
# tokenization
|
||||
embedding_context = nullcontext if self.use_grad else torch.no_grad
|
||||
with embedding_context():
|
||||
x = self.model(tokens.input_ids.to(we.device_id),
|
||||
tokens.attention_mask.to(we.device_id))
|
||||
x = x.last_hidden_state
|
||||
# if not self.return_pooled:
|
||||
# return x.detach()
|
||||
# else:
|
||||
# return x.detach(), self.pool(x, tokens.input_ids)
|
||||
if return_mask:
|
||||
return x.detach() + 0.0, tokens.attention_mask.to(we.device_id)
|
||||
else:
|
||||
return x.detach() + 0.0
|
||||
|
||||
def pool(self, x, tokens):
|
||||
# take features from the eot embedding (eot_token is the highest number in each sequence)
|
||||
return x[torch.arange(x.shape[0]),
|
||||
torch.argmax((tokens.input_ids == 1).float(), dim=-1)]
|
||||
|
||||
def _clean(self, text):
|
||||
if self.clean == 'whitespace':
|
||||
text = whitespace_clean(basic_clean(text))
|
||||
elif self.clean == 'lower':
|
||||
text = whitespace_clean(basic_clean(text)).lower()
|
||||
elif self.clean == 'canonicalize':
|
||||
text = canonicalize(basic_clean(text))
|
||||
elif self.clean == 'heavy':
|
||||
text = heavy_clean(heavy_clean(text))
|
||||
return text
|
||||
|
||||
def encode_text(self,
|
||||
tokens,
|
||||
tokenizer=None,
|
||||
append_sentence_embedding=False,
|
||||
return_mask=False):
|
||||
return self(tokens, return_mask=return_mask)
|
||||
|
||||
def encode(self, text, return_mask=False):
|
||||
if isinstance(text, str):
|
||||
text = [text]
|
||||
if self.clean:
|
||||
text = [self._clean(u) for u in text]
|
||||
assert self.tokenizer is not None
|
||||
tokens = self.tokenizer(text, **self.tokenize_kargs)
|
||||
return self(tokens, return_mask=return_mask)
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
return dict_to_yaml('MODELS',
|
||||
__class__.__name__,
|
||||
T5EmbedderHF.para_dict,
|
||||
set_name=True)
|
||||
|
||||
|
||||
@EMBEDDERS.register_class()
|
||||
class FrozenCLIPEmbedder2(FrozenCLIPEmbedder):
|
||||
"""Uses the CLIP transformer encoder for text (from huggingface)"""
|
||||
para_dict = {
|
||||
'RETURN_POOLED': {
|
||||
'value': False,
|
||||
'description':
|
||||
'Whether return pooled results or not, default False.'
|
||||
}
|
||||
}
|
||||
para_dict.update(FrozenCLIPEmbedder.para_dict)
|
||||
LAYERS = ['hidden', 'last', 'penultimate']
|
||||
|
||||
def __init__(self, cfg, logger=None):
|
||||
super(FrozenCLIPEmbedder, self).__init__(cfg, logger=logger)
|
||||
self.return_pooled = cfg.get('RETURN_POOLED', False)
|
||||
tokenizer_path = cfg.get('TOKENIZER_PATH', None)
|
||||
if tokenizer_path is not None:
|
||||
with FS.get_dir_to_local_dir(tokenizer_path,
|
||||
wait_finish=True) as local_path:
|
||||
self.tokenizer = CLIPTokenizer.from_pretrained(local_path)
|
||||
|
||||
pretrained_model = cfg.get('PRETRAINED_MODEL', None)
|
||||
if pretrained_model is None:
|
||||
raise 'You should set pretrained_model: modelcard.'
|
||||
with FS.get_dir_to_local_dir(cfg.PRETRAINED_MODEL,
|
||||
wait_finish=True) as local_path:
|
||||
self.transformer = CLIPTextModelWithProjection.from_pretrained(
|
||||
local_path)
|
||||
|
||||
self.use_grad = cfg.get('USE_GRAD', False)
|
||||
self.freeze_flag = cfg.get('FREEZE', True)
|
||||
if self.freeze_flag:
|
||||
self.freeze()
|
||||
|
||||
self.max_length = cfg.get('MAX_LENGTH', 77)
|
||||
self.layer = cfg.get('LAYER', 'last')
|
||||
self.layer_idx = cfg.get('LAYER_IDX', None)
|
||||
self.use_final_layer_norm = cfg.get('USE_FINAL_LAYER_NORM', False)
|
||||
assert self.layer in self.LAYERS
|
||||
if self.layer == 'hidden':
|
||||
assert self.layer_idx is not None
|
||||
assert 0 <= abs(self.layer_idx) <= 12
|
||||
|
||||
def _forward(self, text):
|
||||
batch_encoding = self.tokenizer(text,
|
||||
truncation=True,
|
||||
max_length=self.max_length,
|
||||
return_length=True,
|
||||
return_overflowing_tokens=False,
|
||||
padding='max_length',
|
||||
return_tensors='pt')
|
||||
tokens = batch_encoding['input_ids'].to(we.device_id)
|
||||
outputs = self.transformer(input_ids=tokens, output_hidden_states=True)
|
||||
if self.layer == 'last':
|
||||
context = outputs.last_hidden_state
|
||||
elif self.layer == 'penultimate':
|
||||
context = outputs.hidden_states[-2]
|
||||
else:
|
||||
context = outputs.hidden_states[self.layer_idx]
|
||||
|
||||
if self.return_pooled:
|
||||
pooled = outputs[0]
|
||||
return context, pooled
|
||||
return context
|
||||
|
||||
|
||||
@EMBEDDERS.register_class()
|
||||
class SD3TextEmbedder(BaseEmbedder):
|
||||
def __init__(self, cfg, logger=None):
|
||||
super().__init__(cfg, logger=logger)
|
||||
|
||||
clip_l_config = cfg.get('CLIP_L', None)
|
||||
clip_g_config = cfg.get('CLIP_G', None)
|
||||
t5_xxl_config = cfg.get('T5_XXL', None)
|
||||
|
||||
self.clip_l = EMBEDDERS.build(clip_l_config) if clip_l_config else None
|
||||
self.clip_g = EMBEDDERS.build(clip_g_config) if clip_g_config else None
|
||||
self.t5_xxl = EMBEDDERS.build(t5_xxl_config) if t5_xxl_config else None
|
||||
|
||||
self.p_zero = cfg.get('P_ZERO', 0.464)
|
||||
|
||||
def encode(self, text):
|
||||
return self(text)
|
||||
|
||||
def forward(self, text):
|
||||
l_ctx, g_ctx, t5_ctx = None, None, None
|
||||
n = len(text)
|
||||
l_pooled = torch.zeros((n, 768), device=we.device_id)
|
||||
g_pooled = torch.zeros((n, 1280), device=we.device_id)
|
||||
if self.clip_l:
|
||||
with torch.autocast(device_type='cuda',
|
||||
enabled=True,
|
||||
dtype=torch.float16):
|
||||
l_ctx, l_pooled = self.clip_l.encode(text)
|
||||
if self.clip_g:
|
||||
with torch.autocast(device_type='cuda',
|
||||
enabled=True,
|
||||
dtype=torch.float16):
|
||||
g_ctx, g_pooled = self.clip_g.encode(text)
|
||||
if self.t5_xxl:
|
||||
with torch.autocast(device_type='cuda',
|
||||
enabled=True,
|
||||
dtype=torch.float16):
|
||||
t5_ctx = self.t5_xxl.encode(text)
|
||||
|
||||
pooled = torch.cat((l_pooled, g_pooled), dim=-1)
|
||||
|
||||
if l_ctx is not None and g_ctx is not None:
|
||||
lg_ctx = torch.cat([l_ctx, g_ctx], dim=-1)
|
||||
lg_ctx = torch.nn.functional.pad(lg_ctx,
|
||||
(0, 4096 - lg_ctx.shape[-1]))
|
||||
elif l_ctx is not None:
|
||||
lg_ctx = torch.nn.functional.pad(l_ctx,
|
||||
(0, 4096 - l_ctx.shape[-1]))
|
||||
elif g_ctx is not None:
|
||||
lg_ctx = torch.nn.functional.pad(g_ctx, (768, 0))
|
||||
lg_ctx = torch.nn.functional.pad(lg_ctx,
|
||||
(0, 4096 - lg_ctx.shape[-1]))
|
||||
else:
|
||||
lg_ctx = None
|
||||
|
||||
if t5_ctx is not None and lg_ctx is not None:
|
||||
ctx = torch.cat([lg_ctx, t5_ctx], dim=-2)
|
||||
elif t5_ctx is not None:
|
||||
ctx = t5_ctx
|
||||
elif lg_ctx is not None:
|
||||
ctx = lg_ctx
|
||||
else:
|
||||
ctx = torch.zeros((n, 77, 4096), device=we.device_id)
|
||||
|
||||
return ctx, pooled
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
import argparse
|
||||
from scepter.modules.utils.config import Config
|
||||
from scepter.modules.utils.logger import get_logger
|
||||
std_logger = get_logger(name='scepter')
|
||||
parser = argparse.ArgumentParser(description='Argparser for Scepter:\n')
|
||||
cfg = Config(load=True, parser_ins=parser)
|
||||
|
||||
for file_sys in cfg.FILE_SYSTEM:
|
||||
FS.init_fs_client(file_sys)
|
||||
model = SD3TextEmbedder(cfg.COND_STAGE_MODEL,
|
||||
logger=std_logger).to(we.device_id)
|
||||
text = ['a dog is eating food.']
|
||||
ctx, pooled = model(text)
|
||||
print(ctx.shape, pooled.shape)
|
||||
|
||||
@@ -4,4 +4,5 @@ from scepter.modules.model.network.autoencoder import ae_kl
|
||||
from scepter.modules.model.network.classifier import Classifier
|
||||
from scepter.modules.model.network.diffusion import (diffusion, schedules,
|
||||
solvers)
|
||||
from scepter.modules.model.network.ldm import ldm, ldm_sce, ldm_xl
|
||||
from scepter.modules.model.network.ldm import (ldm, ldm_edit, ldm_pixart,
|
||||
ldm_sce, ldm_sd3, ldm_xl)
|
||||
|
||||
@@ -1,10 +1,11 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
|
||||
import re
|
||||
from collections import OrderedDict
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
from scepter.modules.model.network.train_module import TrainModule
|
||||
from scepter.modules.model.registry import BACKBONES, LOSSES, MODELS
|
||||
from scepter.modules.utils.config import dict_to_yaml
|
||||
@@ -87,6 +88,7 @@ class AutoencoderKL(TrainModule):
|
||||
self.pretrained_model = self.cfg.get('PRETRAINED_MODEL', None)
|
||||
self.ignore_keys = self.cfg.get('IGNORE_KEYS', [])
|
||||
self.batch_size = self.cfg.get('BATCH_SIZE', 16)
|
||||
self.use_conv = self.cfg.get('USE_CONV', True)
|
||||
|
||||
self.construct_network()
|
||||
self.init_network()
|
||||
@@ -95,8 +97,12 @@ class AutoencoderKL(TrainModule):
|
||||
z_channels = self.encoder_cfg.Z_CHANNELS
|
||||
self.encoder = BACKBONES.build(self.encoder_cfg, logger=self.logger)
|
||||
self.decoder = BACKBONES.build(self.decoder_cfg, logger=self.logger)
|
||||
self.conv1 = torch.nn.Conv2d(2 * z_channels, 2 * self.embed_dim, 1)
|
||||
self.conv2 = torch.nn.Conv2d(self.embed_dim, z_channels, 1)
|
||||
self.conv1 = torch.nn.Conv2d(
|
||||
2 * z_channels, 2 *
|
||||
self.embed_dim, 1) if self.use_conv else torch.nn.Identity()
|
||||
self.conv2 = torch.nn.Conv2d(
|
||||
self.embed_dim, z_channels,
|
||||
1) if self.use_conv else torch.nn.Identity()
|
||||
|
||||
if self.loss_cfg is not None:
|
||||
self.loss = LOSSES.build(self.loss_cfg, logger=self.logger)
|
||||
@@ -107,7 +113,7 @@ class AutoencoderKL(TrainModule):
|
||||
wait_finish=True) as local_model:
|
||||
self.init_from_ckpt(local_model, ignore_keys=self.ignore_keys)
|
||||
|
||||
def init_from_ckpt(self, path, ignore_keys=list()):
|
||||
def init_from_ckpt(self, path, ignore_keys):
|
||||
if path.find('.safetensors') > -1:
|
||||
from safetensors import safe_open
|
||||
sd = OrderedDict()
|
||||
@@ -122,20 +128,17 @@ class AutoencoderKL(TrainModule):
|
||||
sd = sd['state_dict']
|
||||
|
||||
new_sd = OrderedDict()
|
||||
|
||||
for k, v in sd.items():
|
||||
ignored = False
|
||||
for ik in ignore_keys:
|
||||
if ik in k:
|
||||
if we.rank == 0:
|
||||
self.logger.info(
|
||||
'ignore key {} from state_dict.'.format(k))
|
||||
ignored = True
|
||||
break
|
||||
if self.ignore_keys is not None:
|
||||
if (isinstance(self.ignore_keys, str) and re.match(self.ignore_keys, k)) or \
|
||||
(isinstance(self.ignore_keys, list) and k in self.ignore_keys):
|
||||
continue
|
||||
k = k.replace('post_quant_conv',
|
||||
'conv2') if 'post_quant_conv' in k else k
|
||||
k = k.replace('quant_conv', 'conv1') if 'quant_conv' in k else k
|
||||
if not ignored:
|
||||
new_sd[k] = v
|
||||
k = k.replace('first_stage_model.', '')
|
||||
new_sd[k] = v
|
||||
|
||||
missing, unexpected = self.load_state_dict(new_sd, strict=False)
|
||||
if we.rank == 0:
|
||||
@@ -264,3 +267,14 @@ class AutoencoderKL(TrainModule):
|
||||
__class__.__name__,
|
||||
AutoencoderKL.para_dict,
|
||||
set_name=True)
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
import argparse
|
||||
from scepter.modules.utils.config import Config
|
||||
from scepter.modules.utils.logger import get_logger
|
||||
std_logger = get_logger(name='scepter')
|
||||
parser = argparse.ArgumentParser(description='Argparser for Scepter:\n')
|
||||
cfg = Config(load=True, parser_ins=parser)
|
||||
model = AutoencoderKL(cfg, logger=std_logger)
|
||||
model.load_pretrained_model(cfg.PRETRAINED_MODEL)
|
||||
|
||||
@@ -170,6 +170,45 @@ def adaptive_anisotropic_filter(x, g=None):
|
||||
return y
|
||||
|
||||
|
||||
def extract_into_tensor(a, t, x_shape):
|
||||
b, *_ = t.shape
|
||||
out = a.gather(-1, t)
|
||||
return out.reshape(b, *((1, ) * (len(x_shape) - 1)))
|
||||
|
||||
|
||||
def discretize_timesteps(t_max, t_min, steps, discretization):
|
||||
"""
|
||||
Implementation of timestep discretization methods.
|
||||
"""
|
||||
if discretization == 'leading':
|
||||
steps = torch.arange(t_min, t_max + 1,
|
||||
(t_max - t_min + 1) / steps).flip(0)
|
||||
elif discretization == 'linspace':
|
||||
steps = torch.linspace(t_max, t_min, steps)
|
||||
elif discretization == 'trailing':
|
||||
steps = torch.arange(t_max, t_min - 1, -((t_max - t_min + 1) / steps))
|
||||
else:
|
||||
raise NotImplementedError(
|
||||
f'{discretization} discretization not implemented')
|
||||
return steps.clamp_(t_min, t_max)
|
||||
|
||||
|
||||
def get_scalings_for_boundary_condition(sigma):
|
||||
sigma_data = 0.5
|
||||
c_skip = (1 -
|
||||
sigma**2)**0.5 * sigma_data**2 / (sigma**2 +
|
||||
(1 - sigma**2) * sigma_data**2)
|
||||
c_out = (sigma * sigma_data / (sigma**2 +
|
||||
(1 - sigma**2) * sigma_data**2)**0.5)
|
||||
return c_skip, c_out
|
||||
|
||||
|
||||
def v_to_x0(v, t, x_t, diffusion):
|
||||
sigmas = _i(diffusion.sigmas, t, v)
|
||||
alphas = _i(diffusion.alphas, t, v)
|
||||
return alphas * x_t - sigmas * v
|
||||
|
||||
|
||||
class GaussianDiffusion(object):
|
||||
def __init__(self, sigmas, prediction_type='eps'):
|
||||
assert prediction_type in {'x0', 'eps', 'v'}
|
||||
@@ -666,40 +705,212 @@ class GaussianDiffusion(object):
|
||||
noise)
|
||||
|
||||
|
||||
def extract_into_tensor(a, t, x_shape):
|
||||
b, *_ = t.shape
|
||||
out = a.gather(-1, t)
|
||||
return out.reshape(b, *((1, ) * (len(x_shape) - 1)))
|
||||
class GaussianDiffusionRF(object):
|
||||
def __init__(self, sigmas, prediction_type='rf'):
|
||||
assert prediction_type in {'rf'}
|
||||
self.sigmas = sigmas
|
||||
self.num_timesteps = len(sigmas)
|
||||
|
||||
def diffuse(self, x0, t, noise, sigma):
|
||||
"""
|
||||
Add Gaussian noise to signal x0 according to:
|
||||
q(x_t | x_0) = N(x_t | alpha_t x_0, sigma_t^2 I).
|
||||
"""
|
||||
shape = (x0.size(0), ) + (1, ) * (x0.ndim - 1)
|
||||
sigma = sigma.view(shape)
|
||||
alpha = 1 - sigma
|
||||
xt = alpha * x0 + sigma * noise
|
||||
return xt
|
||||
|
||||
def discretize_timesteps(t_max, t_min, steps, discretization):
|
||||
"""
|
||||
Implementation of timestep discretization methods.
|
||||
"""
|
||||
if discretization == 'leading':
|
||||
steps = torch.arange(t_min, t_max + 1,
|
||||
(t_max - t_min + 1) / steps).flip(0)
|
||||
elif discretization == 'linspace':
|
||||
steps = torch.linspace(t_max, t_min, steps)
|
||||
elif discretization == 'trailing':
|
||||
steps = torch.arange(t_max, t_min - 1, -((t_max - t_min + 1) / steps))
|
||||
else:
|
||||
raise NotImplementedError(
|
||||
f'{discretization} discretization not implemented')
|
||||
return steps.clamp_(t_min, t_max)
|
||||
def denoise(self,
|
||||
xt,
|
||||
t,
|
||||
sigma,
|
||||
model,
|
||||
model_kwargs={},
|
||||
guide_scale=None,
|
||||
guide_rescale=None,
|
||||
cat_uc=False,
|
||||
**kwargs):
|
||||
|
||||
assert sigma is not None
|
||||
shape = (xt.size(0), ) + (1, ) * (xt.ndim - 1)
|
||||
sigma = sigma.view(shape)
|
||||
|
||||
def get_scalings_for_boundary_condition(sigma):
|
||||
sigma_data = 0.5
|
||||
c_skip = (1 -
|
||||
sigma**2)**0.5 * sigma_data**2 / (sigma**2 +
|
||||
(1 - sigma**2) * sigma_data**2)
|
||||
c_out = (sigma * sigma_data / (sigma**2 +
|
||||
(1 - sigma**2) * sigma_data**2)**0.5)
|
||||
return c_skip, c_out
|
||||
# prediction
|
||||
if guide_scale is None:
|
||||
if isinstance(model_kwargs, dict):
|
||||
out = model(xt, t=t, **model_kwargs, **kwargs)
|
||||
elif isinstance(model_kwargs, list) and len(model_kwargs) > 0:
|
||||
out = model(xt, t=t, **model_kwargs[0], **kwargs)
|
||||
else:
|
||||
raise Exception('Error')
|
||||
else:
|
||||
# classifier-free guidance (arXiv:2207.12598)
|
||||
# model_kwargs[0]: conditional kwargs
|
||||
# model_kwargs[1]: non-conditional kwargs
|
||||
assert isinstance(model_kwargs, list) and len(model_kwargs) >= 2
|
||||
if isinstance(guide_scale, float) or isinstance(guide_scale, int):
|
||||
assert len(model_kwargs) == 2
|
||||
if guide_scale == 1.:
|
||||
out = model(xt, t=t, **model_kwargs[0], **kwargs)
|
||||
else:
|
||||
if cat_uc:
|
||||
|
||||
def parse_model_kwargs(prev_value, value):
|
||||
if isinstance(value, torch.Tensor):
|
||||
prev_value = torch.cat([prev_value, value],
|
||||
dim=0)
|
||||
elif isinstance(value, dict):
|
||||
for k, v in value.items():
|
||||
prev_value[k] = parse_model_kwargs(
|
||||
prev_value[k], v)
|
||||
elif isinstance(value, list):
|
||||
for idx, v in enumerate(value):
|
||||
prev_value[idx] = parse_model_kwargs(
|
||||
prev_value[idx], v)
|
||||
return prev_value
|
||||
|
||||
def v_to_x0(v, t, x_t, diffusion):
|
||||
sigmas = _i(diffusion.sigmas, t, v)
|
||||
alphas = _i(diffusion.alphas, t, v)
|
||||
return alphas * x_t - sigmas * v
|
||||
all_model_kwargs = copy.deepcopy(model_kwargs[0])
|
||||
for model_kwarg in model_kwargs[1:]:
|
||||
for key, value in model_kwarg.items():
|
||||
all_model_kwargs[key] = parse_model_kwargs(
|
||||
all_model_kwargs[key], value)
|
||||
all_out = model(xt.repeat(2, 1, 1, 1),
|
||||
t=t.repeat(2),
|
||||
**all_model_kwargs,
|
||||
**kwargs)
|
||||
y_out, u_out = all_out.chunk(2)
|
||||
else:
|
||||
y_out = model(xt, t=t, **model_kwargs[0], **kwargs)
|
||||
u_out = model(xt, t=t, **model_kwargs[1], **kwargs)
|
||||
|
||||
out = u_out + guide_scale * (y_out - u_out)
|
||||
if guide_rescale is not None and guide_rescale > 0.0:
|
||||
assert guide_rescale >= 0 and guide_rescale <= 1
|
||||
ratio = (
|
||||
y_out.flatten(1).std(dim=1) /
|
||||
(out.flatten(1).std(dim=1) + 1e-12)).view((-1, ) + (1, ) *
|
||||
(y_out.ndim - 1))
|
||||
out *= guide_rescale * ratio + (1 - guide_rescale) * 1.0
|
||||
|
||||
x0 = xt - sigma * out
|
||||
return x0
|
||||
|
||||
def loss(self,
|
||||
x0,
|
||||
t,
|
||||
model,
|
||||
model_kwargs={},
|
||||
reduction='mean',
|
||||
noise=None,
|
||||
**kwargs):
|
||||
|
||||
sigma = t / self.num_timesteps
|
||||
shape = (x0.size(0), ) + (1, ) * (x0.ndim - 1)
|
||||
sigma = sigma.view(shape)
|
||||
if noise is None:
|
||||
noise = torch.randn_like(x0)
|
||||
xt = self.diffuse(x0, t, noise, sigma=sigma)
|
||||
out = model(xt, t=t, **model_kwargs, **kwargs)
|
||||
loss = ((xt - sigma * out) - x0)**2
|
||||
# loss = (out - (x0 - noise)) ** 2
|
||||
if reduction == 'mean':
|
||||
loss = loss.flatten(1).mean(dim=1)
|
||||
return loss
|
||||
|
||||
@torch.no_grad()
|
||||
def sample(self,
|
||||
noise,
|
||||
model,
|
||||
model_kwargs={},
|
||||
guide_scale=None,
|
||||
guide_rescale=None,
|
||||
solver='euler',
|
||||
steps=20,
|
||||
shift=3,
|
||||
discretization=None,
|
||||
return_intermediate=None,
|
||||
show_progress=False,
|
||||
seed=-1,
|
||||
intermediate_callback=None,
|
||||
cat_uc=False,
|
||||
**kwargs):
|
||||
# sanity check
|
||||
assert isinstance(steps, (int, torch.LongTensor))
|
||||
assert return_intermediate in (None, 'x0', 'xt')
|
||||
|
||||
# function of diffusion solver
|
||||
solver_fn = {
|
||||
'ddim': sample_ddim,
|
||||
'euler_ancestral': sample_euler_ancestral,
|
||||
'euler': sample_euler,
|
||||
'heun': sample_heun,
|
||||
'dpm2': sample_dpm_2,
|
||||
'dpm2_ancestral': sample_dpm_2_ancestral,
|
||||
'dpmpp_2s_ancestral': sample_dpmpp_2s_ancestral,
|
||||
'dpmpp_2m': sample_dpmpp_2m,
|
||||
'dpmpp_sde': sample_dpmpp_sde,
|
||||
'dpmpp_2m_sde': sample_dpmpp_2m_sde,
|
||||
'dpm2_karras': sample_dpm_2,
|
||||
'dpm2_ancestral_karras': sample_dpm_2_ancestral,
|
||||
'dpmpp_2s_ancestral_karras': sample_dpmpp_2s_ancestral,
|
||||
'dpmpp_2m_karras': sample_dpmpp_2m,
|
||||
'dpmpp_sde_karras': sample_dpmpp_sde,
|
||||
'dpmpp_2m_sde_karras': sample_dpmpp_2m_sde,
|
||||
'onestep': sample_onestep,
|
||||
'multistep': stochastic_iterative_sampler,
|
||||
'multistep2': stochastic_iterative_sampler2,
|
||||
'multistep3': stochastic_iterative_sampler3,
|
||||
'dpmpp_2m_sde_lcm': sample_dpmpp_2m_sde_lcm,
|
||||
}[solver]
|
||||
|
||||
seed = seed if seed >= 0 else random.randint(0, 2**31)
|
||||
intermediates = []
|
||||
|
||||
def model_fn(xt, sigma):
|
||||
# denoising
|
||||
sigma = sigma.repeat(len(xt)).to(xt.device)
|
||||
t = self._sigma_to_t(sigma).round().long()
|
||||
x0 = self.denoise(xt,
|
||||
t,
|
||||
sigma,
|
||||
model,
|
||||
model_kwargs,
|
||||
guide_scale,
|
||||
guide_rescale,
|
||||
cat_uc=cat_uc,
|
||||
**kwargs)
|
||||
|
||||
# collect intermediate outputs
|
||||
if return_intermediate == 'xt':
|
||||
intermediates.append(xt)
|
||||
elif return_intermediate == 'x0':
|
||||
intermediates.append(x0)
|
||||
if intermediate_callback is not None:
|
||||
intermediate_callback(intermediates[-1])
|
||||
return x0
|
||||
|
||||
# get timesteps
|
||||
device = self.sigmas.device
|
||||
sigma_max = self.sigmas[0]
|
||||
sigma_min = self.sigmas[-1]
|
||||
t_max = sigma_max * self.num_timesteps
|
||||
t_min = sigma_min * self.num_timesteps
|
||||
steps = torch.linspace(t_max, t_min, steps).to(device)
|
||||
sigmas = steps / self.num_timesteps
|
||||
sigmas = shift * sigmas / (1 + (shift - 1) * sigmas)
|
||||
sigmas = sigmas.to(torch.float32).to(device)
|
||||
sigmas = torch.cat([sigmas, sigmas.new_zeros([1])])
|
||||
|
||||
kwargs['seed'] = seed
|
||||
# sampling
|
||||
x0 = solver_fn(noise,
|
||||
model_fn,
|
||||
sigmas,
|
||||
show_progress=show_progress,
|
||||
**kwargs)
|
||||
return (x0, intermediates) if return_intermediate is not None else x0
|
||||
|
||||
def _sigma_to_t(self, sigma):
|
||||
return sigma * self.num_timesteps
|
||||
|
||||
@@ -13,6 +13,7 @@ where alpha_t^2 = 1 - sigma_t^2.
|
||||
"""
|
||||
import math
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
__all__ = [
|
||||
@@ -154,6 +155,15 @@ def logsnr_cosine_interp_schedule(n,
|
||||
_logsnr_cosine_interp(n, logsnr_min, logsnr_max, scale_min, scale_max))
|
||||
|
||||
|
||||
def shifted_schedule(n, shift=3):
|
||||
timesteps = np.linspace(1, n, n, dtype=np.float32)[::-1].copy()
|
||||
timesteps = torch.from_numpy(timesteps).to(dtype=torch.float32)
|
||||
|
||||
sigmas = timesteps / n
|
||||
sigmas = shift * sigmas / (1 + (shift - 1) * sigmas)
|
||||
return sigmas
|
||||
|
||||
|
||||
def noise_schedule(schedule='logsnr_cosine_interp',
|
||||
n=1000,
|
||||
zero_terminal_snr=False,
|
||||
@@ -171,7 +181,8 @@ def noise_schedule(schedule='logsnr_cosine_interp',
|
||||
'vp': vp_schedule,
|
||||
'logsnr_cosine': logsnr_cosine_schedule,
|
||||
'logsnr_cosine_shifted': logsnr_cosine_shifted_schedule,
|
||||
'logsnr_cosine_interp': logsnr_cosine_interp_schedule
|
||||
'logsnr_cosine_interp': logsnr_cosine_interp_schedule,
|
||||
'shifted': shifted_schedule,
|
||||
}[schedule](n, **kwargs)
|
||||
|
||||
# post-processing
|
||||
|
||||
@@ -12,9 +12,8 @@ q(x_t | x_0) = N(x_t | alpha_t x_0, sigma_t^2 I),
|
||||
|
||||
where 0 <= sigma_t <= 1 and alpha_t^2 = 1 - sigma_t^2.
|
||||
"""
|
||||
from tqdm.auto import trange
|
||||
|
||||
import torch
|
||||
from tqdm.auto import trange
|
||||
|
||||
__all__ = [
|
||||
'sample_euler', 'sample_euler_ancestral', 'sample_heun', 'sample_dpm_2',
|
||||
@@ -77,8 +76,9 @@ def sample_euler(noise,
|
||||
denoised = model(noise, sigma_hat)
|
||||
x = denoised + sigmas[i + 1] * (gamma + 1) * noise
|
||||
else:
|
||||
_, c_in = get_scalings(sigma_hat)
|
||||
denoised = model(x * c_in, sigma_hat)
|
||||
# _, c_in = get_scalings(sigma_hat)
|
||||
# denoised = model(x * c_in, sigma_hat)
|
||||
denoised = model(x, sigmas[i])
|
||||
d = (x - denoised) / sigma_hat
|
||||
dt = sigmas[i + 1] - sigma_hat
|
||||
x = x + d * dt
|
||||
|
||||
@@ -2,7 +2,9 @@
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
from scepter.modules.model.network.ldm.ldm import LatentDiffusion
|
||||
from scepter.modules.model.network.ldm.ldm_edit import LatentDiffusionEdit
|
||||
from scepter.modules.model.network.ldm.ldm_pixart import LatentDiffusionPixart
|
||||
from scepter.modules.model.network.ldm.ldm_sce import (
|
||||
LatentDiffusionSCEControl, LatentDiffusionSCETuning,
|
||||
LatentDiffusionXLSCEControl, LatentDiffusionXLSCETuning)
|
||||
from scepter.modules.model.network.ldm.ldm_sd3 import LatentDiffusionSD3
|
||||
from scepter.modules.model.network.ldm.ldm_xl import LatentDiffusionXL
|
||||
|
||||
@@ -1,10 +1,12 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import copy
|
||||
import numbers
|
||||
import random
|
||||
from collections import OrderedDict
|
||||
|
||||
import torch
|
||||
|
||||
from scepter.modules.model.network.diffusion.diffusion import GaussianDiffusion
|
||||
from scepter.modules.model.network.diffusion.schedules import noise_schedule
|
||||
from scepter.modules.model.network.train_module import TrainModule
|
||||
@@ -91,8 +93,8 @@ class LatentDiffusion(TrainModule):
|
||||
def init_params(self):
|
||||
self.parameterization = self.cfg.get('PARAMETERIZATION', 'eps')
|
||||
assert self.parameterization in [
|
||||
'eps', 'x0', 'v'
|
||||
], 'currently only supporting "eps" and "x0" and "v"'
|
||||
'eps', 'x0', 'v', 'rf'
|
||||
], 'currently only supporting "eps" and "x0" and "v" and "rf"'
|
||||
self.num_timesteps = self.cfg.get('TIMESTEPS', 1000)
|
||||
|
||||
self.schedule_args = {
|
||||
@@ -137,15 +139,18 @@ class LatentDiffusion(TrainModule):
|
||||
self.default_n_prompt = ''
|
||||
if self.train_n_prompt is None:
|
||||
self.train_n_prompt = ''
|
||||
self.use_ema = self.cfg.get('USE_EMA', True)
|
||||
self.use_ema = self.cfg.get('USE_EMA', False)
|
||||
self.model_ema_config = self.cfg.get('DIFFUSION_MODEL_EMA', None)
|
||||
|
||||
def construct_network(self):
|
||||
self.model = BACKBONES.build(self.model_config, logger=self.logger)
|
||||
self.logger.info('all parameters:{}'.format(count_params(self.model)))
|
||||
if self.use_ema and self.model_ema_config:
|
||||
self.model_ema = BACKBONES.build(self.model_ema_config,
|
||||
logger=self.logger)
|
||||
if self.use_ema:
|
||||
if self.model_ema_config:
|
||||
self.model_ema = BACKBONES.build(self.model_ema_config,
|
||||
logger=self.logger)
|
||||
else:
|
||||
self.model_ema = copy.deepcopy(self.model)
|
||||
self.model_ema = self.model_ema.eval()
|
||||
for param in self.model_ema.parameters():
|
||||
param.requires_grad = False
|
||||
@@ -154,13 +159,15 @@ class LatentDiffusion(TrainModule):
|
||||
if self.tokenizer_config is not None:
|
||||
self.tokenizer = TOKENIZERS.build(self.tokenizer_config,
|
||||
logger=self.logger)
|
||||
|
||||
self.first_stage_model = MODELS.build(self.first_stage_config,
|
||||
logger=self.logger)
|
||||
self.first_stage_model = self.first_stage_model.eval()
|
||||
self.first_stage_model.train = disabled_train
|
||||
for param in self.first_stage_model.parameters():
|
||||
param.requires_grad = False
|
||||
if self.first_stage_config:
|
||||
self.first_stage_model = MODELS.build(self.first_stage_config,
|
||||
logger=self.logger)
|
||||
self.first_stage_model = self.first_stage_model.eval()
|
||||
self.first_stage_model.train = disabled_train
|
||||
for param in self.first_stage_model.parameters():
|
||||
param.requires_grad = False
|
||||
else:
|
||||
self.first_stage_model = None
|
||||
if self.tokenizer_config is not None:
|
||||
self.cond_stage_config.KWARGS = {
|
||||
'vocab_size': self.tokenizer.vocab_size
|
||||
@@ -269,8 +276,8 @@ class LatentDiffusion(TrainModule):
|
||||
ret = {'loss': loss, 'probe_data': {'prompt': prompt}}
|
||||
return ret
|
||||
|
||||
def noise_sample(self, batch_size, h, w, g):
|
||||
noise = torch.empty(batch_size, 4, h, w,
|
||||
def noise_sample(self, batch_size, h, w, g, c=4):
|
||||
noise = torch.empty(batch_size, c, h, w,
|
||||
device=we.device_id).normal_(generator=g)
|
||||
return noise
|
||||
|
||||
|
||||
@@ -0,0 +1,272 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import copy
|
||||
import numbers
|
||||
import random
|
||||
from collections import OrderedDict
|
||||
|
||||
import torch
|
||||
|
||||
from scepter.modules.model.network.diffusion.diffusion import GaussianDiffusion
|
||||
from scepter.modules.model.network.diffusion.schedules import noise_schedule
|
||||
from scepter.modules.model.network.ldm import LatentDiffusion
|
||||
from scepter.modules.model.network.train_module import TrainModule
|
||||
from scepter.modules.model.registry import (BACKBONES, EMBEDDERS, LOSSES,
|
||||
MODELS, TOKENIZERS)
|
||||
from scepter.modules.model.utils.basic_utils import count_params, default
|
||||
from scepter.modules.utils.config import dict_to_yaml
|
||||
from scepter.modules.utils.distribute import we
|
||||
from scepter.modules.utils.file_system import FS
|
||||
|
||||
|
||||
def disabled_train(self, mode=True):
|
||||
"""Overwrite model.train with this function to make sure train/eval mode
|
||||
does not change anymore."""
|
||||
return self
|
||||
|
||||
|
||||
@MODELS.register_class()
|
||||
class LatentDiffusionPixart(LatentDiffusion):
|
||||
para_dict = LatentDiffusion.para_dict
|
||||
para_dict['DECODER_BIAS'] = {'value': 0, 'description': ''}
|
||||
|
||||
def __init__(self, cfg, logger=None):
|
||||
super().__init__(cfg, logger=logger)
|
||||
self.decoder_bias = cfg.get('DECODER_BIAS', 0.5)
|
||||
|
||||
def construct_network(self):
|
||||
self.model = BACKBONES.build(self.model_config, logger=self.logger)
|
||||
self.logger.info('all parameters:{}'.format(count_params(self.model)))
|
||||
if self.use_ema:
|
||||
self.model_ema = copy.deepcopy(self.model).eval()
|
||||
for param in self.model_ema.parameters():
|
||||
param.requires_grad = False
|
||||
if self.loss_config:
|
||||
self.loss = LOSSES.build(self.loss_config, logger=self.logger)
|
||||
if self.tokenizer_config is not None:
|
||||
self.tokenizer = TOKENIZERS.build(self.tokenizer_config,
|
||||
logger=self.logger)
|
||||
|
||||
if self.first_stage_config:
|
||||
self.first_stage_model = MODELS.build(self.first_stage_config,
|
||||
logger=self.logger)
|
||||
self.first_stage_model = self.first_stage_model.eval()
|
||||
self.first_stage_model.train = disabled_train
|
||||
for param in self.first_stage_model.parameters():
|
||||
param.requires_grad = False
|
||||
else:
|
||||
self.first_stage_model = None
|
||||
if self.tokenizer_config is not None:
|
||||
self.cond_stage_config.KWARGS = {
|
||||
'vocab_size': self.tokenizer.vocab_size
|
||||
}
|
||||
if self.cond_stage_config == '__is_unconditional__':
|
||||
print(
|
||||
f'Training {self.__class__.__name__} as an unconditional model.'
|
||||
)
|
||||
self.cond_stage_model = None
|
||||
else:
|
||||
model = EMBEDDERS.build(self.cond_stage_config, logger=self.logger)
|
||||
self.cond_stage_model = model.eval().requires_grad_(False)
|
||||
self.cond_stage_model.train = disabled_train
|
||||
|
||||
def forward_train(self,
|
||||
image=None,
|
||||
noise=None,
|
||||
prompt=None,
|
||||
label=None,
|
||||
**kwargs):
|
||||
n, c, h, w = image.shape
|
||||
x_start = self.encode_first_stage(image, **kwargs)
|
||||
t = torch.randint(0, self.num_timesteps, (n, ),
|
||||
device=x_start.device).long()
|
||||
ar = torch.tensor([[h / w]], device=we.device_id).repeat(n, 1)
|
||||
hw = torch.tensor([[h, w]], dtype=torch.float,
|
||||
device=we.device_id).repeat(n, 1)
|
||||
context = {}
|
||||
cont_mask = None
|
||||
if prompt and self.cond_stage_model:
|
||||
with torch.autocast(device_type='cuda',
|
||||
enabled=True,
|
||||
dtype=torch.bfloat16):
|
||||
cont, cont_mask = getattr(self.cond_stage_model,
|
||||
'encode')(prompt, return_mask=True)
|
||||
context['crossattn'] = cont.float()
|
||||
else:
|
||||
assert label is not None
|
||||
context['label'] = label
|
||||
|
||||
if 'hint' in kwargs and kwargs['hint'] is not None:
|
||||
hint = kwargs.pop('hint')
|
||||
if isinstance(context, dict):
|
||||
context['hint'] = hint
|
||||
else:
|
||||
context = {'crossattn': context, 'hint': hint}
|
||||
else:
|
||||
hint = None
|
||||
if self.min_snr_gamma is not None:
|
||||
alphas = self.diffusion.alphas.to(we.device_id)[t]
|
||||
sigmas = self.diffusion.sigmas.pow(2).to(we.device_id)[t]
|
||||
snrs = (alphas / sigmas).clamp(min=1e-20)
|
||||
min_snrs = snrs.clamp(max=self.min_snr_gamma)
|
||||
weights = min_snrs / snrs
|
||||
else:
|
||||
weights = 1
|
||||
self.register_probe({'snrs_weights': weights})
|
||||
|
||||
loss = self.diffusion.loss(x0=x_start,
|
||||
t=t,
|
||||
model=self.model,
|
||||
model_kwargs={
|
||||
'cond': context,
|
||||
'mask': cont_mask,
|
||||
'data_info': {
|
||||
'img_hw': hw,
|
||||
'aspect_ratio': ar
|
||||
}
|
||||
},
|
||||
noise=noise,
|
||||
**kwargs)
|
||||
loss = loss * weights
|
||||
loss = loss.mean()
|
||||
ret = {'loss': loss, 'probe_data': {'prompt': prompt}}
|
||||
return ret
|
||||
|
||||
@torch.no_grad()
|
||||
def forward_test(self,
|
||||
prompt=None,
|
||||
label=None,
|
||||
sampler='ddim',
|
||||
sample_steps=20,
|
||||
seed=2023,
|
||||
guide_scale=4.5,
|
||||
guide_rescale=0.5,
|
||||
discretization='trailing',
|
||||
**kwargs):
|
||||
g = torch.Generator(device=we.device_id)
|
||||
seed = seed if seed >= 0 else random.randint(0, 2**32 - 1)
|
||||
g.manual_seed(seed)
|
||||
num_samples = label.shape[0] if label is not None else len(prompt)
|
||||
context = {}
|
||||
null_context = {}
|
||||
cont_mask = None
|
||||
if prompt and self.cond_stage_model:
|
||||
with torch.autocast(device_type='cuda',
|
||||
enabled=True,
|
||||
dtype=torch.bfloat16):
|
||||
cont, cont_mask = getattr(self.cond_stage_model,
|
||||
'encode')(prompt, return_mask=True)
|
||||
context['crossattn'] = cont.float()
|
||||
null_context['crossattn'] = self.model.y_embedder.y_embedding[
|
||||
None].repeat(len(prompt), 1, 1)
|
||||
else:
|
||||
assert label is not None
|
||||
context['label'] = label
|
||||
null_context['label'] = torch.tensor(
|
||||
[self.model.num_classes]).repeat(num_samples).to(we.device_id)
|
||||
|
||||
if 'hint' in kwargs and kwargs['hint'] is not None:
|
||||
hint = kwargs.pop('hint')
|
||||
if isinstance(context, dict):
|
||||
context['hint'] = hint
|
||||
if isinstance(null_context, dict):
|
||||
null_context['hint'] = hint
|
||||
else:
|
||||
hint = None
|
||||
if 'index' in kwargs:
|
||||
kwargs.pop('index')
|
||||
image_size = None
|
||||
if 'meta' in kwargs:
|
||||
meta = kwargs.pop('meta')
|
||||
if 'image_size' in meta:
|
||||
h = int(meta['image_size'][0][0])
|
||||
w = int(meta['image_size'][1][0])
|
||||
image_size = [h, w]
|
||||
if 'image_size' in kwargs:
|
||||
image_size = kwargs.pop('image_size')
|
||||
if isinstance(image_size, numbers.Number):
|
||||
image_size = [image_size, image_size]
|
||||
if image_size is None:
|
||||
image_size = [1024, 1024]
|
||||
height, width = image_size
|
||||
noise = self.noise_sample(num_samples, height // self.size_factor,
|
||||
width // self.size_factor, g)
|
||||
# UNet use input n_prompt
|
||||
samples = self.diffusion.sample(
|
||||
solver=sampler,
|
||||
noise=noise,
|
||||
model=self.model,
|
||||
model_kwargs=[{
|
||||
'cond': context,
|
||||
'mask': cont_mask,
|
||||
'data_info': {
|
||||
'img_hw':
|
||||
torch.tensor([image_size],
|
||||
dtype=torch.float,
|
||||
device=we.device_id).repeat(num_samples, 1),
|
||||
'aspect_ratio':
|
||||
torch.tensor([[1.]], device=we.device_id).repeat(
|
||||
num_samples, 1)
|
||||
}
|
||||
}, {
|
||||
'cond': null_context,
|
||||
'mask': cont_mask,
|
||||
'data_info': {
|
||||
'img_hw':
|
||||
torch.tensor([image_size],
|
||||
dtype=torch.float,
|
||||
device=we.device_id).repeat(num_samples, 1),
|
||||
'aspect_ratio':
|
||||
torch.tensor([[1.]], device=we.device_id).repeat(
|
||||
num_samples, 1)
|
||||
}
|
||||
}] if guide_scale is not None and guide_scale > 0 else {
|
||||
'cond': context,
|
||||
'mask': cont_mask,
|
||||
'data_info': {
|
||||
'img_hw':
|
||||
torch.tensor([image_size],
|
||||
dtype=torch.float,
|
||||
device=we.device_id).repeat(num_samples, 1),
|
||||
'aspect_ratio':
|
||||
torch.tensor([[1.]], device=we.device_id).repeat(
|
||||
num_samples, 1)
|
||||
}
|
||||
},
|
||||
cat_uc=False,
|
||||
steps=sample_steps,
|
||||
guide_scale=guide_scale,
|
||||
guide_rescale=guide_rescale,
|
||||
discretization=discretization,
|
||||
show_progress=True,
|
||||
seed=seed,
|
||||
condition_fn=None,
|
||||
clamp=None,
|
||||
percentile=None,
|
||||
t_max=None,
|
||||
t_min=None,
|
||||
discard_penultimate_step=None,
|
||||
return_intermediate=None,
|
||||
**kwargs)
|
||||
x_samples = self.decode_first_stage(samples).float()
|
||||
x_samples = torch.clamp(
|
||||
(x_samples + 1.0) / 2.0 + self.decoder_bias / 255,
|
||||
min=0.0,
|
||||
max=1.0)
|
||||
outputs = list()
|
||||
prompt = label.detach().cpu().numpy().tolist(
|
||||
) if prompt is None else prompt
|
||||
for i, (p, img) in enumerate(zip(prompt, x_samples)):
|
||||
one_tup = {'prompt': str(p), 'n_prompt': '', 'image': img}
|
||||
if hint is not None:
|
||||
one_tup.update({'hint': hint[i]})
|
||||
outputs.append(one_tup)
|
||||
|
||||
return outputs
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
return dict_to_yaml('MODEL',
|
||||
__class__.__name__,
|
||||
LatentDiffusionPixart.para_dict,
|
||||
set_name=True)
|
||||
@@ -0,0 +1,236 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import numbers
|
||||
import random
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from scepter.modules.model.network.diffusion.diffusion import \
|
||||
GaussianDiffusionRF
|
||||
from scepter.modules.model.network.diffusion.schedules import noise_schedule
|
||||
from scepter.modules.model.network.ldm import LatentDiffusion
|
||||
from scepter.modules.model.registry import MODELS
|
||||
from scepter.modules.utils.config import dict_to_yaml
|
||||
from scepter.modules.utils.distribute import we
|
||||
|
||||
|
||||
@MODELS.register_class()
|
||||
class LatentDiffusionSD3(LatentDiffusion):
|
||||
para_dict = LatentDiffusion.para_dict
|
||||
|
||||
def __init__(self, cfg, logger=None):
|
||||
super().__init__(cfg, logger=logger)
|
||||
|
||||
self.shift_factor = cfg.get('SHIFT_FACTOR', 0)
|
||||
self.t_weight_type = cfg.get('T_WEIGHT', 'logit_normal')
|
||||
self.logit_mean = cfg.get('LOGIT_MEAN', 0.0)
|
||||
self.logit_std = cfg.get('LOGIT_STD', 1.0)
|
||||
|
||||
def init_params(self):
|
||||
self.parameterization = self.cfg.get('PARAMETERIZATION', 'rf')
|
||||
assert self.parameterization in [
|
||||
'eps', 'x0', 'v', 'rf'
|
||||
], 'currently only supporting "eps" and "x0" and "v" and "rf"'
|
||||
self.num_timesteps = self.cfg.get('TIMESTEPS', 1000)
|
||||
|
||||
self.schedule_args = {
|
||||
k.lower(): v
|
||||
for k, v in self.cfg.get('SCHEDULE_ARGS', {
|
||||
'NAME': 'logsnr_cosine_interp',
|
||||
'SCALE_MIN': 2.0,
|
||||
'SCALE_MAX': 4.0
|
||||
}).items()
|
||||
}
|
||||
|
||||
self.min_snr_gamma = self.cfg.get('MIN_SNR_GAMMA', None)
|
||||
|
||||
self.zero_terminal_snr = self.cfg.get('ZERO_TERMINAL_SNR', False)
|
||||
if self.zero_terminal_snr:
|
||||
assert self.parameterization == 'v', 'Now zero_terminal_snr only support v-prediction mode.'
|
||||
|
||||
self.sigmas = noise_schedule(schedule=self.schedule_args.pop('name'),
|
||||
n=self.num_timesteps,
|
||||
zero_terminal_snr=self.zero_terminal_snr,
|
||||
**self.schedule_args)
|
||||
|
||||
self.diffusion = GaussianDiffusionRF(
|
||||
sigmas=self.sigmas, prediction_type=self.parameterization)
|
||||
|
||||
self.pretrained_model = self.cfg.get('PRETRAINED_MODEL', None)
|
||||
self.ignore_keys = self.cfg.get('IGNORE_KEYS', [])
|
||||
|
||||
self.model_config = self.cfg.DIFFUSION_MODEL
|
||||
self.first_stage_config = self.cfg.FIRST_STAGE_MODEL
|
||||
self.cond_stage_config = self.cfg.COND_STAGE_MODEL
|
||||
self.tokenizer_config = self.cfg.get('TOKENIZER', None)
|
||||
self.loss_config = self.cfg.get('LOSS', None)
|
||||
|
||||
self.scale_factor = self.cfg.get('SCALE_FACTOR', 0.18215)
|
||||
self.size_factor = self.cfg.get('SIZE_FACTOR', 8)
|
||||
self.default_n_prompt = self.cfg.get('DEFAULT_N_PROMPT', '')
|
||||
self.default_n_prompt = '' if self.default_n_prompt is None else self.default_n_prompt
|
||||
self.p_zero = self.cfg.get('P_ZERO', 0.0)
|
||||
self.train_n_prompt = self.cfg.get('TRAIN_N_PROMPT', '')
|
||||
if self.default_n_prompt is None:
|
||||
self.default_n_prompt = ''
|
||||
if self.train_n_prompt is None:
|
||||
self.train_n_prompt = ''
|
||||
self.use_ema = self.cfg.get('USE_EMA', False)
|
||||
self.model_ema_config = self.cfg.get('DIFFUSION_MODEL_EMA', None)
|
||||
|
||||
def noise_sample(self, batch_size, h, w, g, c=4):
|
||||
noise = torch.empty(batch_size, c, h, w,
|
||||
device=we.device_id).normal_(generator=g)
|
||||
return noise
|
||||
|
||||
def forward_train(self, image=None, noise=None, prompt=None, **kwargs):
|
||||
n, c, h, w = image.shape
|
||||
x_start = self.encode_first_stage(image, **kwargs)
|
||||
if self.t_weight_type == 'uniform':
|
||||
t = torch.randint(0,
|
||||
self.num_timesteps, (n, ),
|
||||
device=x_start.device).long()
|
||||
elif self.t_weight_type == 'logit_normal':
|
||||
density = F.sigmoid(
|
||||
torch.normal(mean=self.logit_mean,
|
||||
std=self.logit_std,
|
||||
size=(n, ),
|
||||
device=x_start.device))
|
||||
t = (density * (self.num_timesteps - 1)).round().long()
|
||||
sigma = (t + 1) / self.num_timesteps
|
||||
shift = self.schedule_args['shift']
|
||||
if shift > 1.:
|
||||
sigma = shift * sigma / (1 + (shift - 1) * sigma)
|
||||
t = sigma * self.num_timesteps
|
||||
|
||||
context = {}
|
||||
if prompt and self.cond_stage_model:
|
||||
ctx, pooled = getattr(self.cond_stage_model, 'encode')(prompt)
|
||||
context['crossattn'] = ctx.float()
|
||||
context['y'] = pooled
|
||||
else:
|
||||
assert False
|
||||
|
||||
if self.min_snr_gamma is not None:
|
||||
alphas = self.diffusion.alphas.to(we.device_id)[t]
|
||||
sigmas = self.diffusion.sigmas.pow(2).to(we.device_id)[t]
|
||||
snrs = (alphas / sigmas).clamp(min=1e-20)
|
||||
min_snrs = snrs.clamp(max=self.min_snr_gamma)
|
||||
weights = min_snrs / snrs
|
||||
else:
|
||||
weights = 1
|
||||
self.register_probe({'snrs_weights': weights})
|
||||
|
||||
loss = self.diffusion.loss(x0=x_start,
|
||||
t=t,
|
||||
model=self.model,
|
||||
model_kwargs={
|
||||
'cond': context,
|
||||
},
|
||||
noise=noise,
|
||||
**kwargs)
|
||||
loss = loss * weights
|
||||
loss = loss.mean()
|
||||
ret = {'loss': loss, 'probe_data': {'prompt': prompt}}
|
||||
return ret
|
||||
|
||||
@torch.no_grad()
|
||||
def forward_test(self,
|
||||
prompt=None,
|
||||
sampler='ddim',
|
||||
sample_steps=20,
|
||||
seed=2023,
|
||||
guide_scale=4.5,
|
||||
guide_rescale=0.0,
|
||||
**kwargs):
|
||||
g = torch.Generator(device=we.device_id)
|
||||
seed = seed if seed >= 0 else random.randint(0, 2**32 - 1)
|
||||
g.manual_seed(seed)
|
||||
num_samples = len(prompt)
|
||||
context = {}
|
||||
null_context = {}
|
||||
|
||||
if prompt and self.cond_stage_model:
|
||||
ctx, pooled = getattr(self.cond_stage_model, 'encode')(prompt)
|
||||
null_ctx, null_pooled = getattr(self.cond_stage_model,
|
||||
'encode')([''] * len(prompt))
|
||||
context['crossattn'] = ctx.float()
|
||||
context['y'] = pooled.float()
|
||||
null_context['crossattn'] = null_ctx.float()
|
||||
null_context['y'] = null_pooled.float()
|
||||
else:
|
||||
assert False
|
||||
|
||||
if 'index' in kwargs:
|
||||
kwargs.pop('index')
|
||||
image_size = None
|
||||
if 'meta' in kwargs:
|
||||
meta = kwargs.pop('meta')
|
||||
if 'image_size' in meta:
|
||||
h = int(meta['image_size'][0][0])
|
||||
w = int(meta['image_size'][1][0])
|
||||
image_size = [h, w]
|
||||
if 'image_size' in kwargs:
|
||||
image_size = kwargs.pop('image_size')
|
||||
if isinstance(image_size, numbers.Number):
|
||||
image_size = [image_size, image_size]
|
||||
if image_size is None:
|
||||
image_size = [1024, 1024]
|
||||
height, width = image_size
|
||||
noise = self.noise_sample(num_samples,
|
||||
height // self.size_factor,
|
||||
width // self.size_factor,
|
||||
g,
|
||||
c=16)
|
||||
# UNet use input n_prompt
|
||||
samples = self.diffusion.sample(
|
||||
solver=sampler,
|
||||
noise=noise,
|
||||
model=self.model,
|
||||
model_kwargs=[{
|
||||
'cond': context
|
||||
}, {
|
||||
'cond': null_context
|
||||
}] if guide_scale is not None and guide_scale > 0 else {
|
||||
'cond': context,
|
||||
},
|
||||
cat_uc=False,
|
||||
steps=sample_steps,
|
||||
guide_scale=guide_scale,
|
||||
guide_rescale=guide_rescale,
|
||||
show_progress=True,
|
||||
seed=seed,
|
||||
condition_fn=None,
|
||||
clamp=None,
|
||||
percentile=None,
|
||||
t_max=None,
|
||||
t_min=None,
|
||||
discard_penultimate_step=None,
|
||||
return_intermediate=None,
|
||||
**kwargs)
|
||||
x_samples = self.decode_first_stage(samples).float()
|
||||
x_samples = torch.clamp((x_samples + 1.0) / 2.0, min=0.0, max=1.0)
|
||||
outputs = list()
|
||||
|
||||
for i, (p, img) in enumerate(zip(prompt, x_samples)):
|
||||
one_tup = {'prompt': str(p), 'n_prompt': '', 'image': img}
|
||||
outputs.append(one_tup)
|
||||
|
||||
return outputs
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
return dict_to_yaml('MODEL',
|
||||
__class__.__name__,
|
||||
LatentDiffusionSD3.para_dict,
|
||||
set_name=True)
|
||||
|
||||
@torch.no_grad()
|
||||
def encode_first_stage(self, x, **kwargs):
|
||||
z = self.first_stage_model.encode(x)
|
||||
return self.scale_factor * (z - self.shift_factor)
|
||||
|
||||
@torch.no_grad()
|
||||
def decode_first_stage(self, z):
|
||||
z = 1. / self.scale_factor * z + self.shift_factor
|
||||
return self.first_stage_model.decode(z)
|
||||
@@ -4,7 +4,7 @@ import open_clip
|
||||
from scepter.modules.model.registry import TOKENIZERS
|
||||
from scepter.modules.model.tokenizer import BaseTokenizer
|
||||
from scepter.modules.model.tokenizer.tokenizer_component import (
|
||||
basic_clean, whitespace_clean)
|
||||
basic_clean, canonicalize, heavy_clean, whitespace_clean)
|
||||
from scepter.modules.utils.config import dict_to_yaml
|
||||
from scepter.modules.utils.file_system import FS
|
||||
from transformers import CLIPTokenizer as transformer_clip_tokenizer
|
||||
|
||||
@@ -1,10 +1,13 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import html
|
||||
import string
|
||||
from functools import lru_cache
|
||||
from urllib import parse
|
||||
|
||||
import ftfy
|
||||
import regex as re
|
||||
from bs4 import BeautifulSoup
|
||||
|
||||
|
||||
@lru_cache()
|
||||
@@ -59,3 +62,132 @@ def whitespace_clean(text):
|
||||
text = re.sub(r'\s+', ' ', text)
|
||||
text = text.strip()
|
||||
return text
|
||||
|
||||
|
||||
def canonicalize(text, keep_punctuation_exact_string=None):
|
||||
text = text.replace('_', ' ')
|
||||
if keep_punctuation_exact_string:
|
||||
text = keep_punctuation_exact_string.join(
|
||||
part.translate(str.maketrans('', '', string.punctuation))
|
||||
for part in text.split(keep_punctuation_exact_string))
|
||||
else:
|
||||
text = text.translate(str.maketrans('', '', string.punctuation))
|
||||
text = text.lower()
|
||||
text = re.sub(r'\s+', ' ', text)
|
||||
return text.strip()
|
||||
|
||||
|
||||
def heavy_clean(text):
|
||||
text = str(text)
|
||||
text = parse.unquote_plus(text)
|
||||
text = text.strip().lower()
|
||||
text = re.sub('<person>', 'person', text)
|
||||
|
||||
# urls:
|
||||
text = re.sub(
|
||||
r'\b((?:https?:(?:\/{1,3}|[a-zA-Z0-9%])|[a-zA-Z0-9.\-]+[.](?:com|co|ru|net|org|edu|gov|it)[\w/-]*\b\/?(?!@)))', # noqa: E501
|
||||
'',
|
||||
text)
|
||||
text = re.sub(
|
||||
r'\b((?:www:(?:\/{1,3}|[a-zA-Z0-9%])|[a-zA-Z0-9.\-]+[.](?:com|co|ru|net|org|edu|gov|it)[\w/-]*\b\/?(?!@)))', # noqa: E501
|
||||
'',
|
||||
text)
|
||||
|
||||
# html:
|
||||
text = BeautifulSoup(text, features='html.parser').text
|
||||
|
||||
# @<nickname>
|
||||
text = re.sub(r'@[\w\d]+\b', '', text)
|
||||
|
||||
# 31C0—31EF CJK Strokes
|
||||
# 31F0—31FF Katakana Phonetic Extensions
|
||||
# 3200—32FF Enclosed CJK Letters and Months
|
||||
# 3300—33FF CJK Compatibility
|
||||
# 3400—4DBF CJK Unified Ideographs Extension A
|
||||
# 4DC0—4DFF Yijing Hexagram Symbols
|
||||
# 4E00—9FFF CJK Unified Ideographs
|
||||
text = re.sub(r'[\u31c0-\u31ef]+', '', text)
|
||||
text = re.sub(r'[\u31f0-\u31ff]+', '', text)
|
||||
text = re.sub(r'[\u3200-\u32ff]+', '', text)
|
||||
text = re.sub(r'[\u3300-\u33ff]+', '', text)
|
||||
text = re.sub(r'[\u3400-\u4dbf]+', '', text)
|
||||
text = re.sub(r'[\u4dc0-\u4dff]+', '', text)
|
||||
text = re.sub(r'[\u4e00-\u9fff]+', '', text)
|
||||
#######################################################
|
||||
|
||||
# все виды тире / all types of dash --> "-"
|
||||
text = re.sub(
|
||||
r'[\u002D\u058A\u05BE\u1400\u1806\u2010-\u2015\u2E17\u2E1A\u2E3A\u2E3B\u2E40\u301C\u3030\u30A0\uFE31\uFE32\uFE58\uFE63\uFF0D]+', # noqa: E501
|
||||
'-',
|
||||
text)
|
||||
|
||||
# кавычки к одному стандарту
|
||||
text = re.sub(r'[`´«»“”¨]', '"', text)
|
||||
text = re.sub(r'[‘’]', "'", text)
|
||||
|
||||
# "
|
||||
text = re.sub(r'"?', '', text)
|
||||
# &
|
||||
text = re.sub(r'&', '', text)
|
||||
|
||||
# ip adresses:
|
||||
text = re.sub(r'\d{1,3}\.\d{1,3}\.\d{1,3}\.\d{1,3}', ' ', text)
|
||||
|
||||
# article ids:
|
||||
text = re.sub(r'\d:\d\d\s+$', '', text)
|
||||
|
||||
# \n
|
||||
text = re.sub(r'\\n', ' ', text)
|
||||
|
||||
# "#123"
|
||||
text = re.sub(r'#\d{1,3}\b', '', text)
|
||||
# "#12345.."
|
||||
text = re.sub(r'#\d{5,}\b', '', text)
|
||||
# "123456.."
|
||||
text = re.sub(r'\b\d{6,}\b', '', text)
|
||||
# filenames:
|
||||
text = re.sub(r'[\S]+\.(?:png|jpg|jpeg|bmp|webp|eps|pdf|apk|mp4)', '',
|
||||
text)
|
||||
|
||||
text = re.sub(r'[\"\']{2,}', r'"', text) # """AUSVERKAUFT"""
|
||||
text = re.sub(r'[\.]{2,}', r' ', text) # """AUSVERKAUFT"""
|
||||
text = re.sub(
|
||||
re.compile(r'[' + '#®•©™&@·º½¾¿¡§~' + '\)' + '\(' + '\]' + # noqa
|
||||
'\[' + # noqa
|
||||
'\}' + '\{' + '\|' + '\\' + '\/' + '\*' + # noqa
|
||||
r']{1,}'), # noqa
|
||||
r' ',
|
||||
text) # ***AUSVERKAUFT***, #AUSVERKAUFT
|
||||
text = re.sub(r'\s+\.\s+', r' ', text) # " . "
|
||||
|
||||
# this-is-my-cute-cat / this_is_my_cute_cat
|
||||
regex2 = re.compile(r'(?:\-|\_)')
|
||||
if len(re.findall(regex2, text)) > 3:
|
||||
text = re.sub(regex2, ' ', text)
|
||||
|
||||
text = basic_clean(text)
|
||||
|
||||
text = re.sub(r'\b[a-zA-Z]{1,3}\d{3,15}\b', '', text) # jc6640
|
||||
text = re.sub(r'\b[a-zA-Z]+\d+[a-zA-Z]+\b', '', text) # jc6640vc
|
||||
text = re.sub(r'\b\d+[a-zA-Z]+\d+\b', '', text) # 6640vc231
|
||||
|
||||
text = re.sub(r'(worldwide\s+)?(free\s+)?shipping', '', text)
|
||||
text = re.sub(r'(free\s)?download(\sfree)?', '', text)
|
||||
text = re.sub(r'\bclick\b\s(?:for|on)\s\w+', '', text)
|
||||
text = re.sub(r'\b(?:png|jpg|jpeg|bmp|webp|eps|pdf|apk|mp4)(\simage[s]?)?',
|
||||
'', text)
|
||||
text = re.sub(r'\bpage\s+\d+\b', '', text)
|
||||
|
||||
# j2d1a2a...
|
||||
text = re.sub(r'\b\d*[a-zA-Z]+\d+[a-zA-Z]+\d+[a-zA-Z\d]*\b', r' ', text)
|
||||
|
||||
text = re.sub(r'\b\d+\.?\d*[xх×]\d+\.?\d*\b', '', text)
|
||||
text = re.sub(r'\b\s+\:\s+', r': ', text)
|
||||
text = re.sub(r'(\D[,\./])\b', r'\1 ', text)
|
||||
text = re.sub(r'\s+', ' ', text)
|
||||
|
||||
text = re.sub(r'^[\"\']([\w\W]+)[\"\']$', r'\1', text)
|
||||
text = re.sub(r'^[\'\_,\-\:;]', r'', text)
|
||||
text = re.sub(r'[\'\_,\-\:\-\+]$', r'', text)
|
||||
text = re.sub(r'^\.\S+$', '', text)
|
||||
return text.strip()
|
||||
|
||||
@@ -4,8 +4,6 @@ import copy
|
||||
import os
|
||||
from collections import OrderedDict, defaultdict
|
||||
|
||||
from tqdm import tqdm
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.cuda.amp as amp
|
||||
@@ -23,6 +21,7 @@ from torch.distributed.fsdp import (BackwardPrefetch, CPUOffload,
|
||||
FullyShardedDataParallel, MixedPrecision,
|
||||
ShardingStrategy, StateDictType)
|
||||
from torch.nn.parallel import DistributedDataParallel
|
||||
from tqdm import tqdm
|
||||
|
||||
sharding_strategy_map = {
|
||||
'full_shard': ShardingStrategy.FULL_SHARD,
|
||||
@@ -343,9 +342,7 @@ class LatentDiffusionSolver(BaseSolver):
|
||||
self.train_mode()
|
||||
batch_data = next(data_iter)
|
||||
self.before_iter(self.hooks_dict[self._mode])
|
||||
if self.sample_args:
|
||||
batch_data.update(self.sample_args.get_lowercase_dict())
|
||||
if 'meta' in batch_data:
|
||||
if 'meta' in batch_data and isinstance(batch_data['meta'], dict):
|
||||
self.register_probe({
|
||||
'data_key':
|
||||
ProbeData(batch_data['meta'].get('data_key', []),
|
||||
@@ -356,6 +353,8 @@ class LatentDiffusionSolver(BaseSolver):
|
||||
'batch_size': len(batch_data['prompt'])
|
||||
})
|
||||
self.current_batch_data[self.mode] = batch_data
|
||||
if self.sample_args:
|
||||
batch_data.update(self.sample_args.get_lowercase_dict())
|
||||
with torch.autocast(device_type='cuda',
|
||||
enabled=self.use_amp,
|
||||
dtype=self.dtype):
|
||||
@@ -524,6 +523,8 @@ class LatentDiffusionSolver(BaseSolver):
|
||||
if len(swift_cfg_dict) > 0:
|
||||
from swift import Swift
|
||||
model = Swift.prepare_model(self.model, config=swift_cfg_dict)
|
||||
|
||||
# self.logger.info([(key, param.shape) for key, param in self.model.named_parameters() if param.requires_grad])
|
||||
return model
|
||||
|
||||
def freeze(self, freeze_cfg, model=None):
|
||||
|
||||
@@ -132,8 +132,9 @@ class CheckpointHook(Hook):
|
||||
else:
|
||||
if hasattr(solver, 'save_pretrained'):
|
||||
save_path = osp.join(
|
||||
solver.work_dir, 'checkpoints/{}-{}-bin'.format(
|
||||
self.save_name_prefix, solver.total_iter + 1))
|
||||
solver.work_dir,
|
||||
'checkpoints/{}-{}'.format(self.save_name_prefix,
|
||||
solver.total_iter + 1))
|
||||
local_folder, _ = FS.map_to_local(save_path)
|
||||
FS.make_dir(local_folder)
|
||||
ckpt, cfg = solver.save_pretrained()
|
||||
|
||||
@@ -134,11 +134,11 @@ def load_pretrained_dict(model: torch.nn.Module,
|
||||
missing_keys = load_status.missing_keys
|
||||
err_msgs = []
|
||||
if unexpected_keys:
|
||||
err_msgs.append('unexpected key in source '
|
||||
f'state_dict: {", ".join(unexpected_keys)}\n')
|
||||
err_msgs.append('unexpected key in source state_dict: {}\n'.format(
|
||||
', '.join(unexpected_keys)))
|
||||
if missing_keys:
|
||||
err_msgs.append('missing key in source '
|
||||
f'state_dict: {", ".join(missing_keys)}\n')
|
||||
err_msgs.append('missing key in source state_dict: {}\n'.format(
|
||||
', '.join(missing_keys)))
|
||||
err_msgs = '\n'.join(err_msgs)
|
||||
|
||||
if len(err_msgs) > 0:
|
||||
|
||||
@@ -2,6 +2,8 @@
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
from scepter.modules.inference.diffusion_inference import DiffusionInference
|
||||
from scepter.modules.inference.largen_inference import LargenInference
|
||||
from scepter.modules.inference.pixart_inference import PixArtInference
|
||||
from scepter.modules.inference.sd3_inference import SD3Inference
|
||||
from scepter.modules.inference.stylebooth_inference import StyleboothInference
|
||||
from scepter.modules.utils.logger import get_logger
|
||||
|
||||
@@ -102,6 +104,10 @@ class PipelineManager():
|
||||
PipelineBuilder = LargenInference
|
||||
elif pipeline_name.startswith('EDIT'):
|
||||
PipelineBuilder = StyleboothInference
|
||||
elif pipeline_name.startswith('PIXART'):
|
||||
PipelineBuilder = PixArtInference
|
||||
elif pipeline_name.startswith('SD3'):
|
||||
PipelineBuilder = SD3Inference
|
||||
else:
|
||||
PipelineBuilder = DiffusionInference
|
||||
new_inference = PipelineBuilder(logger=self.logger)
|
||||
|
||||
@@ -6,9 +6,15 @@ from scepter.modules.utils.file_system import FS
|
||||
|
||||
|
||||
def download_image(image):
|
||||
# return None
|
||||
if image is not None:
|
||||
name = get_md5(image)
|
||||
local_path = FS.get_from(image, f'/tmp/gradio/scepter_examples/{name}')
|
||||
client = FS.get_fs_client(image)
|
||||
if client.tmp_dir.startswith('/home'):
|
||||
name = get_md5(image)
|
||||
local_path = FS.get_from(image,
|
||||
f'/tmp/gradio/scepter_examples/{name}')
|
||||
else:
|
||||
local_path = FS.get_from(image)
|
||||
return local_path
|
||||
else:
|
||||
return image
|
||||
@@ -48,20 +54,20 @@ class ModelManageUIName():
|
||||
if language == 'en':
|
||||
self.model_block_name = 'Model Management'
|
||||
self.postprocess_model_name = 'Refiners and Tuners'
|
||||
self.diffusion_model = 'Unet'
|
||||
self.diffusion_model = 'Model'
|
||||
self.first_stage_model = 'Vae'
|
||||
self.cond_stage_model = 'Condition Model'
|
||||
self.refine_diffusion_model = 'Refine Unet'
|
||||
self.refine_diffusion_model = 'Refine Model'
|
||||
self.refine_cond_model = 'Refine Condition Model'
|
||||
self.load_lora_tuner = 'Load Lora Tuner'
|
||||
self.load_swift_tuner = 'Load swift Tuner'
|
||||
elif language == 'zh':
|
||||
self.model_block_name = '模型管理'
|
||||
self.postprocess_model_name = 'Refiners and Tuners'
|
||||
self.diffusion_model = 'Unet'
|
||||
self.diffusion_model = 'Model'
|
||||
self.first_stage_model = 'Vae'
|
||||
self.cond_stage_model = 'Condition Model'
|
||||
self.refine_diffusion_model = 'Refine Unet'
|
||||
self.refine_diffusion_model = 'Refine Model'
|
||||
self.refine_cond_model = 'Refine Condition Model'
|
||||
self.load_lora_tuner = '加载 Lora 微调模型'
|
||||
self.load_swift_tuner = '加载 Swift 微调模型'
|
||||
@@ -371,6 +377,7 @@ class LargenUIName():
|
||||
self.masked_image = 'Masked Image'
|
||||
self.preprocess = 'Input Preprocess'
|
||||
self.button_name = 'Data Preprocess'
|
||||
self.value = 'CenterAround'
|
||||
self.direction = (
|
||||
'Instruction: \n\n'
|
||||
'For customized data: \n\n'
|
||||
@@ -404,6 +411,7 @@ class LargenUIName():
|
||||
self.masked_image = '掩码图片'
|
||||
self.preprocess = '输入图片预处理'
|
||||
self.button_name = '数据预处理'
|
||||
self.value = '中心向外'
|
||||
self.direction = ('使用说明:\n'
|
||||
'针对自定义数据:\n'
|
||||
'a.1)上传背景图片(待编辑);\n'
|
||||
@@ -435,7 +443,7 @@ class LargenUIName():
|
||||
None,
|
||||
1.0,
|
||||
0.75,
|
||||
'CenterAround',
|
||||
self.value,
|
||||
1024,
|
||||
1024
|
||||
],
|
||||
@@ -456,7 +464,7 @@ class LargenUIName():
|
||||
),
|
||||
1.0,
|
||||
0.0,
|
||||
'',
|
||||
None,
|
||||
1024,
|
||||
1024
|
||||
],
|
||||
@@ -473,7 +481,7 @@ class LargenUIName():
|
||||
None,
|
||||
1.0,
|
||||
0.0,
|
||||
'',
|
||||
None,
|
||||
1024,
|
||||
1024
|
||||
],
|
||||
@@ -494,7 +502,7 @@ class LargenUIName():
|
||||
),
|
||||
0.45,
|
||||
0.0,
|
||||
'',
|
||||
None,
|
||||
1024,
|
||||
1024
|
||||
],
|
||||
|
||||
@@ -75,13 +75,13 @@ class ControlUI(UIBase):
|
||||
self.source_image = gr.Image(
|
||||
label=self.component_names.source_image,
|
||||
type='pil',
|
||||
tool='editor',
|
||||
sources=['upload'],
|
||||
interactive=True)
|
||||
with gr.Column(scale=1, min_width=0):
|
||||
self.cond_image = gr.Image(
|
||||
label=self.component_names.cond_image,
|
||||
type='pil',
|
||||
tool='editor',
|
||||
sources=['upload'],
|
||||
interactive=True)
|
||||
with gr.Row():
|
||||
with gr.Column(scale=1, min_width=0):
|
||||
|
||||
@@ -47,12 +47,10 @@ class GalleryUI(UIBase):
|
||||
container=False,
|
||||
autofocus=True,
|
||||
elem_classes='type_row',
|
||||
submit_on_enter=True,
|
||||
lines=1)
|
||||
|
||||
with gr.Column(scale=3, min_width=0):
|
||||
self.generate_button = gr.Button(
|
||||
label='Generate',
|
||||
value=self.component_names.generate,
|
||||
elem_classes='type_row',
|
||||
elem_id='generate_button',
|
||||
@@ -280,7 +278,7 @@ class GalleryUI(UIBase):
|
||||
**kwargs):
|
||||
self.manager = manager
|
||||
self.gen_inputs = list(self.component_mapping.values())
|
||||
print(self.gen_inputs, len(self.gen_inputs))
|
||||
# print(self.gen_inputs, len(self.gen_inputs))
|
||||
self.gen_outputs = [
|
||||
self.before_refine_panel,
|
||||
self.before_refine_gallery,
|
||||
|
||||
@@ -68,30 +68,29 @@ class LargenUI(UIBase):
|
||||
with gr.Column(scale=2, min_width=0):
|
||||
with gr.Row():
|
||||
with gr.Column(scale=1, min_width=0):
|
||||
self.scene_image = gr.Image(
|
||||
self.scene_image = gr.ImageMask(
|
||||
label=self.component_names.scene_image,
|
||||
type='pil',
|
||||
tool='sketch',
|
||||
source='upload',
|
||||
height=400,
|
||||
sources=['upload'],
|
||||
layers=False,
|
||||
interactive=True)
|
||||
self.cache_button = gr.Button(
|
||||
value='Use Last Generated Image', visible=True)
|
||||
with gr.Column(scale=1, min_width=0):
|
||||
self.subject_image = gr.Image(
|
||||
self.subject_image = gr.ImageMask(
|
||||
label=self.component_names.subject_image,
|
||||
type='pil',
|
||||
tool='sketch',
|
||||
source='upload',
|
||||
interactive=True,
|
||||
height=400,
|
||||
visible=False)
|
||||
sources=['upload'],
|
||||
layers=False,
|
||||
visible=False,
|
||||
interactive=True)
|
||||
|
||||
self.gallery = gr.Gallery(label='Image History',
|
||||
value=[],
|
||||
columns=1,
|
||||
rows=1,
|
||||
height=500)
|
||||
height=500,
|
||||
interactive=False)
|
||||
self.clear_button = gr.Button(value='Clear History',
|
||||
visible=True)
|
||||
|
||||
@@ -255,7 +254,7 @@ class LargenUI(UIBase):
|
||||
queue=False)
|
||||
|
||||
def read_gallery_image(gallery):
|
||||
if len(gallery) == 0:
|
||||
if gallery is None:
|
||||
last_image = None
|
||||
else:
|
||||
last_image = gallery[-1]['name']
|
||||
@@ -267,7 +266,7 @@ class LargenUI(UIBase):
|
||||
|
||||
def clear_gallery(image_history, gallery):
|
||||
image_history.clear()
|
||||
gallery.clear()
|
||||
gallery = []
|
||||
return image_history, gallery
|
||||
|
||||
self.clear_button.click(fn=clear_gallery,
|
||||
@@ -276,8 +275,8 @@ class LargenUI(UIBase):
|
||||
|
||||
def data_process(scene_image, subject_image, task, image_ratio,
|
||||
out_direction, output_height, output_width):
|
||||
tar_image = scene_image['image'].convert('RGB')
|
||||
tar_mask = scene_image['mask'].convert('L')
|
||||
tar_image = scene_image['background'].convert('RGB')
|
||||
tar_mask = scene_image['layers'][0].split()[-1].convert('L')
|
||||
tar_image = np.asarray(tar_image)
|
||||
tar_mask = np.asarray(tar_mask)
|
||||
tar_mask = np.where(tar_mask > 128, 1, 0).astype(np.uint8)
|
||||
@@ -288,8 +287,8 @@ class LargenUI(UIBase):
|
||||
output_height,
|
||||
output_width)
|
||||
elif task == 'Subject_Guided_Inpainting':
|
||||
ref_image = subject_image['image'].convert('RGB')
|
||||
ref_mask = subject_image['mask'].convert('L')
|
||||
ref_image = subject_image['background'].convert('RGB')
|
||||
ref_mask = subject_image['layers'][0].split()[-1].convert('L')
|
||||
ref_image = np.asarray(ref_image)
|
||||
ref_mask = np.asarray(ref_mask)
|
||||
ref_mask = np.where(ref_mask > 128, 1, 0).astype(np.uint8)
|
||||
@@ -298,8 +297,8 @@ class LargenUI(UIBase):
|
||||
1.3, output_height,
|
||||
output_width)
|
||||
elif task == 'Text_Subject_Guided_Inpainting':
|
||||
ref_image = subject_image['image'].convert('RGB')
|
||||
ref_mask = subject_image['mask'].convert('L')
|
||||
ref_image = subject_image['background'].convert('RGB')
|
||||
ref_mask = subject_image['layers'][0].split()[-1].convert('L')
|
||||
ref_image = np.asarray(ref_image)
|
||||
ref_mask = np.asarray(ref_mask)
|
||||
ref_mask = np.where(ref_mask > 128, 1, 0).astype(np.uint8)
|
||||
|
||||
@@ -47,8 +47,8 @@ class MantraUI(UIBase):
|
||||
|
||||
def create_ui(self, *args, **kwargs):
|
||||
self.state = gr.State(value=False)
|
||||
with gr.Column(equal_height=True, visible=False) as self.tab:
|
||||
with gr.Row(scale=1):
|
||||
with gr.Column(visible=False) as self.tab:
|
||||
with gr.Row():
|
||||
with gr.Column(scale=1):
|
||||
with gr.Group(visible=True):
|
||||
with gr.Row(equal_height=True):
|
||||
@@ -99,7 +99,7 @@ class MantraUI(UIBase):
|
||||
with gr.Row(equal_height=True):
|
||||
self.style_example = gr.Image(
|
||||
label=self.component_names.style_example,
|
||||
source='upload',
|
||||
sources=['upload'],
|
||||
value=None,
|
||||
interactive=False)
|
||||
with gr.Row(equal_height=True):
|
||||
|
||||
@@ -142,6 +142,7 @@ class ModelManageUI(UIBase):
|
||||
last_pipeline_ins = self.pipe_manager.pipeline_level_modules[
|
||||
last_pipline]
|
||||
last_pipeline_ins.dynamic_unload(name='all')
|
||||
|
||||
now_pipeline = self.pipe_manager.model_level_info[diffusion_model][
|
||||
'pipeline'][0]
|
||||
pipeline_ins = self.pipe_manager.pipeline_level_modules[
|
||||
@@ -202,9 +203,11 @@ class ModelManageUI(UIBase):
|
||||
value=controller_default),
|
||||
gr.Dropdown(choices=mantra_ui.all_styles.get(now_pipeline, []),
|
||||
value=[]),
|
||||
gr.Textbox(choices=cur_paras.NEGATIVE_PROMPT.get('VALUES', []),
|
||||
gr.Textbox(elem_classes=cur_paras.NEGATIVE_PROMPT.get(
|
||||
'VALUES', []),
|
||||
value=cur_paras.NEGATIVE_PROMPT.get('DEFAULT', '')),
|
||||
gr.Textbox(choices=cur_paras.PROMPT_PREFIX.get('VALUES', []),
|
||||
gr.Textbox(elem_classes=cur_paras.PROMPT_PREFIX.get(
|
||||
'VALUES', []),
|
||||
value=cur_paras.PROMPT_PREFIX.get('DEFAULT', '')),
|
||||
gr.Dropdown(choices=[key for key in h_level_dict.keys()],
|
||||
value=default_res[0]),
|
||||
|
||||
@@ -15,7 +15,7 @@ class StyleboothUI(UIBase):
|
||||
|
||||
def create_ui(self, *args, **kwargs):
|
||||
self.state = gr.State(value=False)
|
||||
with gr.Column(equal_height=False, visible=False) as self.tab:
|
||||
with gr.Column(visible=False) as self.tab:
|
||||
with gr.Row():
|
||||
self.selected_app = gr.Dropdown(
|
||||
label=self.component_names.dropdown_name,
|
||||
@@ -30,7 +30,7 @@ class StyleboothUI(UIBase):
|
||||
self.edit_image = gr.Image(
|
||||
label=self.component_names.source_image,
|
||||
type='pil',
|
||||
tool='editor',
|
||||
sources=['upload'],
|
||||
interactive=True)
|
||||
with gr.Column(variant='panel',
|
||||
scale=1,
|
||||
@@ -58,7 +58,6 @@ class StyleboothUI(UIBase):
|
||||
interactive=True,
|
||||
allow_custom_value=True)
|
||||
self.compose_instruction = gr.Button(
|
||||
label=self.component_names.compose_button,
|
||||
value=self.component_names.compose_button,
|
||||
elem_classes='type_row',
|
||||
elem_id='push')
|
||||
|
||||
@@ -51,8 +51,8 @@ class TunerUI(UIBase):
|
||||
|
||||
def create_ui(self, *args, **kwargs):
|
||||
self.state = gr.State(value=False)
|
||||
with gr.Column(equal_height=True, visible=False) as self.tab:
|
||||
with gr.Row(scale=1):
|
||||
with gr.Column(visible=False) as self.tab:
|
||||
with gr.Row():
|
||||
with gr.Column(variant='panel', scale=1, min_width=0):
|
||||
with gr.Group(visible=True):
|
||||
with gr.Row(equal_height=True):
|
||||
@@ -72,7 +72,6 @@ class TunerUI(UIBase):
|
||||
multiselect=True,
|
||||
interactive=True)
|
||||
self.save_button = gr.Button(
|
||||
label=self.component_names.save_button,
|
||||
value=self.component_names.save_button,
|
||||
elem_classes='type_row',
|
||||
elem_id='save_button',
|
||||
@@ -96,7 +95,7 @@ class TunerUI(UIBase):
|
||||
with gr.Row(equal_height=True):
|
||||
self.tuner_example = gr.Image(
|
||||
label=self.component_names.tuner_example,
|
||||
source='upload',
|
||||
sources=['upload'],
|
||||
value=None,
|
||||
interactive=False)
|
||||
with gr.Row(equal_height=True):
|
||||
|
||||
@@ -3,12 +3,11 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import datetime
|
||||
import json
|
||||
import os.path
|
||||
from collections import OrderedDict
|
||||
|
||||
import gradio as gr
|
||||
from tqdm import tqdm
|
||||
|
||||
from scepter.modules.utils.file_system import FS
|
||||
from scepter.studio.preprocess.caption_editor_ui.component_names import \
|
||||
CreateDatasetUIName
|
||||
@@ -17,6 +16,7 @@ from scepter.studio.preprocess.utils.img2img_data_card import \
|
||||
from scepter.studio.preprocess.utils.txt2img_data_card import \
|
||||
Text2ImageDataCard
|
||||
from scepter.studio.utils.uibase import UIBase
|
||||
from tqdm import tqdm
|
||||
|
||||
refresh_symbol = '\U0001f504' # 🔄
|
||||
|
||||
@@ -144,10 +144,17 @@ class CreateDatasetUI(UIBase):
|
||||
meta_file = os.path.join(one_dir, 'meta.json')
|
||||
if not FS.exists(meta_file):
|
||||
continue
|
||||
with FS.get_from(meta_file) as local_meta:
|
||||
meta = json.load(open(local_meta, 'r'))
|
||||
dataset_type_name = 'scepter_' + meta.get('task_type', '')
|
||||
dataset_type = self.default_dataset_type
|
||||
dataset_cls = self.default_dataset_cls
|
||||
for key, value in self.dataset_type_dict.items():
|
||||
if one_dir.split('/')[-1].startswith(key):
|
||||
if one_dir.endswith('/'):
|
||||
folder_name = one_dir[:-1].split('/')[-1]
|
||||
else:
|
||||
folder_name = one_dir.split('/')[-1]
|
||||
if folder_name.startswith(key) or dataset_type_name == key:
|
||||
dataset_type = key
|
||||
dataset_cls = value
|
||||
break
|
||||
@@ -164,9 +171,9 @@ class CreateDatasetUI(UIBase):
|
||||
return dataset_list
|
||||
|
||||
def create_ui(self):
|
||||
with gr.Box():
|
||||
with gr.Group():
|
||||
gr.Markdown(self.components_name.user_direction)
|
||||
with gr.Box():
|
||||
with gr.Group():
|
||||
with gr.Row():
|
||||
self.sys_log = gr.Markdown(
|
||||
self.components_name.system_log.format(''))
|
||||
@@ -518,8 +525,8 @@ class CreateDatasetUI(UIBase):
|
||||
if prev_data_name in dataset_list:
|
||||
dataset_list.remove(prev_data_name)
|
||||
dataset_list.append(user_dataset_name)
|
||||
self.user_level_dataset_list[
|
||||
login_user_name] = dataset_list
|
||||
self.user_level_dataset_list[login_user_name][
|
||||
dataset_type] = dataset_list
|
||||
self.dataset_dict[dataset_type].pop(prev_data_name)
|
||||
self.dataset_dict[dataset_type][
|
||||
user_dataset_name] = ori_dataset_ins
|
||||
@@ -582,8 +589,11 @@ class CreateDatasetUI(UIBase):
|
||||
user_level_dataset_list = self.load_history(
|
||||
login_user_name=login_user_name)
|
||||
self.user_level_dataset_list.update(user_level_dataset_list)
|
||||
dataset_list = user_level_dataset_list[login_user_name].get(
|
||||
trans_dataset_type, [])
|
||||
if user_level_dataset_list.get(login_user_name) is not None:
|
||||
dataset_list = user_level_dataset_list[login_user_name].get(
|
||||
trans_dataset_type, [])
|
||||
else:
|
||||
dataset_list = []
|
||||
return gr.Dropdown(
|
||||
value=dataset_list[-1] if len(dataset_list) > 0 else '',
|
||||
choices=dataset_list), self.components_name.system_log.format(
|
||||
|
||||
@@ -39,13 +39,13 @@ class DatasetGalleryUI(UIBase):
|
||||
self.default_image_format = os.path.splitext(
|
||||
current_info.get('relative_path', ''))[-1]
|
||||
|
||||
self.default_edit_image_width = gr.Text(value=current_info.get(
|
||||
'edit_width', current_info.get('width', -1)))
|
||||
self.default_edit_image_height = gr.Text(value=current_info.get(
|
||||
'edit_height', current_info.get('height', -1)))
|
||||
self.default_edit_image_format = gr.Text(value=os.path.splitext(
|
||||
self.default_edit_image_width = current_info.get(
|
||||
'edit_width', current_info.get('width', -1))
|
||||
self.default_edit_image_height = current_info.get(
|
||||
'edit_height', current_info.get('height', -1))
|
||||
self.default_edit_image_format = os.path.splitext(
|
||||
current_info.get('edit_relative_path',
|
||||
current_info.get('relative_path', '')))[-1])
|
||||
current_info.get('relative_path', '')))[-1]
|
||||
|
||||
self.default_select_index = self.default_dataset.cursor
|
||||
self.default_info = f'{self.default_dataset.cursor + 1}/{len(self.default_dataset)}'
|
||||
@@ -120,6 +120,8 @@ class DatasetGalleryUI(UIBase):
|
||||
elem_id='dataset_tag_editor_dataset_gallery',
|
||||
value=self.default_image_list,
|
||||
selected_index=self.default_select_index,
|
||||
object_fit='fill',
|
||||
preview=True,
|
||||
columns=4,
|
||||
visible=False,
|
||||
interactive=False)
|
||||
@@ -162,9 +164,7 @@ class DatasetGalleryUI(UIBase):
|
||||
preview=True,
|
||||
allow_preview=True,
|
||||
object_fit='fill')
|
||||
# with gr.Column(variant='panel', scale=2, min_width=0):
|
||||
|
||||
# with gr.Row(visible=False) as :
|
||||
with gr.Row(visible=False) as self.edit_setting_panel:
|
||||
self.sys_log = gr.Markdown(
|
||||
self.component_names.system_log.format(''))
|
||||
@@ -173,7 +173,7 @@ class DatasetGalleryUI(UIBase):
|
||||
visible=False,
|
||||
scale=1,
|
||||
min_width=0) as self.edit_confirm_panel:
|
||||
with gr.Box():
|
||||
with gr.Group():
|
||||
with gr.Row():
|
||||
gr.Markdown(self.component_names.confirm_direction)
|
||||
with gr.Row():
|
||||
@@ -235,6 +235,7 @@ class DatasetGalleryUI(UIBase):
|
||||
with gr.Column(variant='panel', scale=1, min_width=0):
|
||||
self.upload_image = gr.Image(
|
||||
label=self.component_names.upload_image,
|
||||
sources=['upload'],
|
||||
type='pil')
|
||||
with gr.Row():
|
||||
with gr.Column(min_width=0):
|
||||
@@ -251,7 +252,7 @@ class DatasetGalleryUI(UIBase):
|
||||
variant='panel',
|
||||
visible=False,
|
||||
) as self.image_preprocess_panel:
|
||||
with gr.Box():
|
||||
with gr.Group():
|
||||
with gr.Column(variant='panel', min_width=0):
|
||||
with gr.Row():
|
||||
self.image_preprocess_method = gr.Dropdown(
|
||||
@@ -294,7 +295,7 @@ class DatasetGalleryUI(UIBase):
|
||||
image_preprocess_btn)
|
||||
with gr.Row(variant='panel',
|
||||
visible=False) as self.caption_preprocess_panel:
|
||||
with gr.Box():
|
||||
with gr.Group():
|
||||
with gr.Column(variant='panel', min_width=0):
|
||||
with gr.Row():
|
||||
self.caption_preprocess_method = gr.Dropdown(
|
||||
|
||||
@@ -112,7 +112,11 @@ class TrainerUIName():
|
||||
self.train_data = 'Training Data'
|
||||
self.eval_prompts = 'Eval Prompts'
|
||||
self.eval_image = 'Eval Image'
|
||||
self.training_block = 'Training Parameters'
|
||||
self.training_block = '''
|
||||
### Training Parameters
|
||||
'''
|
||||
self.model_param = 'Model Parameters'
|
||||
self.base_param = 'Base Parameters'
|
||||
self.base_model = 'Base Model'
|
||||
self.tuner_name = 'Tuner Method'
|
||||
self.base_model_revision = 'Model Version Number'
|
||||
@@ -175,7 +179,11 @@ class TrainerUIName():
|
||||
self.train_data = '训练数据'
|
||||
self.eval_prompts = '评测文本'
|
||||
self.eval_image = '评测图片'
|
||||
self.training_block = '训练参数'
|
||||
self.training_block = '''
|
||||
### 训练参数
|
||||
'''
|
||||
self.model_param = '模型参数'
|
||||
self.base_param = '基本参数'
|
||||
self.base_model = '基础模型'
|
||||
self.tuner_name = '微调方法'
|
||||
self.base_model_revision = '模型版本号'
|
||||
|
||||
@@ -89,21 +89,21 @@ class ModelUI(UIBase):
|
||||
return all_gallery_list
|
||||
|
||||
def create_ui(self, *args, **kwargs):
|
||||
with gr.Box():
|
||||
with gr.Group():
|
||||
gr.Markdown(self.component_names.output_model_block)
|
||||
with gr.Row(variant='panel', equal_height=True):
|
||||
with gr.Column(scale=7, min_width=0):
|
||||
self.output_model_name = gr.Dropdown(
|
||||
label=self.component_names.output_model_name,
|
||||
choices=self.model_list,
|
||||
value=self.base_model_info.get('model_default', ''),
|
||||
value=None,
|
||||
show_label=False,
|
||||
container=False,
|
||||
interactive=True)
|
||||
with gr.Column(scale=2, min_width=0):
|
||||
self.output_ckpt_name = gr.Dropdown(
|
||||
label=self.component_names.output_ckpt_name,
|
||||
value='',
|
||||
value=None,
|
||||
choices=[],
|
||||
show_label=False,
|
||||
container=False,
|
||||
@@ -127,7 +127,7 @@ class ModelUI(UIBase):
|
||||
self.confirm_add = gr.Button(value=confirm_symbol)
|
||||
with gr.Row(variant='panel', equal_height=True):
|
||||
with gr.Column(scale=3, min_width=0):
|
||||
with gr.Box():
|
||||
with gr.Group():
|
||||
self.log_message = gr.Text(
|
||||
placeholder='Please select model'
|
||||
'or press export button.',
|
||||
|
||||
@@ -8,10 +8,9 @@ import random
|
||||
from collections import OrderedDict
|
||||
|
||||
import gradio as gr
|
||||
import yaml
|
||||
|
||||
import scepter
|
||||
import torch
|
||||
import yaml
|
||||
from scepter.modules.utils.file_system import FS
|
||||
from scepter.studio.self_train.scripts.trainer import TrainManager
|
||||
from scepter.studio.self_train.self_train_ui.component_names import \
|
||||
@@ -98,27 +97,25 @@ class TrainerUI(UIBase):
|
||||
reverse=False))
|
||||
|
||||
def create_ui(self):
|
||||
with gr.Box():
|
||||
with gr.Tabs():
|
||||
with gr.Row(variant='panel', equal_height=True):
|
||||
with gr.Column(scale=1, min_width=0, variant='panel'):
|
||||
gr.Markdown(self.component_names.user_direction)
|
||||
with gr.Box():
|
||||
with gr.Group():
|
||||
gr.Markdown(self.component_names.train_data)
|
||||
self.data_source = gr.Dropdown(
|
||||
choices=self.component_names.data_source_choices,
|
||||
value=self.component_names.data_source_value,
|
||||
label=self.component_names.data_source_name,
|
||||
placeholder=self.component_names.data_source_name,
|
||||
interactive=False)
|
||||
self.data_type = gr.Dropdown(
|
||||
choices=[
|
||||
self.component_names.data_type_map[key]
|
||||
for key in self.component_names.data_type_choices
|
||||
self.component_names.data_type_map[key] for key
|
||||
in self.component_names.data_type_choices
|
||||
],
|
||||
value=self.component_names.data_type_map[
|
||||
self.component_names.data_type_value],
|
||||
label=self.component_names.data_type_name,
|
||||
placeholder=self.component_names.data_type_name,
|
||||
interactive=False)
|
||||
self.ori_data_name = gr.Textbox(
|
||||
label=self.component_names.ori_data_name,
|
||||
@@ -133,7 +130,7 @@ class TrainerUI(UIBase):
|
||||
ms_data_name_place_hold,
|
||||
visible=False,
|
||||
interactive=False)
|
||||
with gr.Box(visible=False) as self.ms_data_box:
|
||||
with gr.Group(visible=False) as self.ms_data_box:
|
||||
with gr.Row():
|
||||
self.ms_data_space = gr.Textbox(
|
||||
label=self.component_names.ms_data_space,
|
||||
@@ -142,7 +139,7 @@ class TrainerUI(UIBase):
|
||||
label=self.component_names.ms_data_subname,
|
||||
value='default',
|
||||
max_lines=1)
|
||||
with gr.Box():
|
||||
with gr.Group():
|
||||
gr.Markdown(self.component_names.eval_data)
|
||||
self.eval_prompts = gr.Dropdown(
|
||||
value=None,
|
||||
@@ -161,18 +158,18 @@ class TrainerUI(UIBase):
|
||||
self.eval_image = gr.Image(
|
||||
label=self.component_names.eval_image,
|
||||
type='pil',
|
||||
tool='editor',
|
||||
sources=['upload'],
|
||||
interactive=False,
|
||||
visible=False)
|
||||
with gr.Column(scale=2, min_width=0, variant='panel'):
|
||||
with gr.Box():
|
||||
gr.Markdown(self.component_names.training_block)
|
||||
gr.Markdown(self.component_names.training_block)
|
||||
with gr.Group():
|
||||
gr.Markdown(self.component_names.model_param)
|
||||
with gr.Row():
|
||||
self.base_model = gr.Dropdown(
|
||||
choices=self.para_data.get(
|
||||
'model_choices', []),
|
||||
value=self.para_data.get(
|
||||
'model_default', ''),
|
||||
value=self.para_data.get('model_default', ''),
|
||||
label=self.component_names.base_model,
|
||||
interactive=True)
|
||||
self.base_model_revision = gr.Dropdown(
|
||||
@@ -180,114 +177,114 @@ class TrainerUI(UIBase):
|
||||
'version_choices', []),
|
||||
value=self.para_data.get(
|
||||
'version_default', ''),
|
||||
label=self.component_names.
|
||||
base_model_revision,
|
||||
label=self.component_names.base_model_revision,
|
||||
interactive=True)
|
||||
with gr.Box():
|
||||
gr.Markdown(self.component_names.tuner_param)
|
||||
# with gr.Column(scale=1, min_width=0):
|
||||
if 'TUNER' in self.para_data:
|
||||
lora_visible, text_lora_visible, sce_visible = judge_tuner_visible(
|
||||
self.para_data['TUNER'])
|
||||
else:
|
||||
lora_visible, text_lora_visible, sce_visible = False, False, False
|
||||
with gr.Row():
|
||||
self.tuner_name = gr.Dropdown(
|
||||
choices=self.para_data.get(
|
||||
'tuner_choices', []),
|
||||
with gr.Group():
|
||||
gr.Markdown(self.component_names.tuner_param)
|
||||
# with gr.Column(scale=1, min_width=0):
|
||||
if 'TUNER' in self.para_data:
|
||||
lora_visible, text_lora_visible, sce_visible = judge_tuner_visible(
|
||||
self.para_data['TUNER'])
|
||||
else:
|
||||
lora_visible, text_lora_visible, sce_visible = False, False, False
|
||||
with gr.Row():
|
||||
self.tuner_name = gr.Dropdown(
|
||||
choices=self.para_data.get(
|
||||
'tuner_choices', []),
|
||||
value=self.para_data.get('tuner_default', ''),
|
||||
label=self.component_names.tuner_name,
|
||||
interactive=True)
|
||||
with gr.Row():
|
||||
with gr.Row(
|
||||
visible=lora_visible) as self.lora_param:
|
||||
self.lora_alpha = gr.Number(
|
||||
label='LoRA Alpha',
|
||||
value=self.para_data.get(
|
||||
'tuner_default', ''),
|
||||
label=self.component_names.tuner_name,
|
||||
'lora_alpha', 256),
|
||||
interactive=True)
|
||||
self.lora_rank = gr.Number(
|
||||
label='LoRA Rank',
|
||||
value=self.para_data.get('lora_rank', 256),
|
||||
interactive=True)
|
||||
with gr.Row(visible=text_lora_visible
|
||||
) as self.text_lora_param:
|
||||
self.text_lora_alpha = gr.Number(
|
||||
label='Text LoRA Alpha',
|
||||
value=self.para_data.get(
|
||||
'text_lora_alpha', 256),
|
||||
interactive=True)
|
||||
self.text_lora_rank = gr.Number(
|
||||
label='Text LoRA Rank',
|
||||
value=self.para_data.get(
|
||||
'text_lora_rank', 256),
|
||||
interactive=True)
|
||||
with gr.Row():
|
||||
with gr.Row(visible=lora_visible) as self.lora_param:
|
||||
self.lora_alpha = gr.Number(
|
||||
label='LoRA Alpha',
|
||||
value=self.para_data.get(
|
||||
'lora_alpha', 256),
|
||||
interactive=True)
|
||||
self.lora_rank = gr.Number(
|
||||
label='LoRA Rank',
|
||||
value=self.para_data.get(
|
||||
'lora_rank', 256),
|
||||
interactive=True)
|
||||
with gr.Row(visible=text_lora_visible) as self.text_lora_param:
|
||||
self.text_lora_alpha = gr.Number(
|
||||
label='Text LoRA Alpha',
|
||||
value=self.para_data.get(
|
||||
'text_lora_alpha', 256),
|
||||
interactive=True)
|
||||
self.text_lora_rank = gr.Number(
|
||||
label='Text LoRA Rank',
|
||||
value=self.para_data.get(
|
||||
'text_lora_rank', 256),
|
||||
interactive=True)
|
||||
self.sce_ratio = gr.Slider(
|
||||
label='SCE Ratio',
|
||||
minimum=0.2,
|
||||
maximum=2.0,
|
||||
step=0.1,
|
||||
value=self.para_data.get(
|
||||
'sce_ratio', 1.0),
|
||||
value=self.para_data.get('sce_ratio', 1.0),
|
||||
interactive=True,
|
||||
visible=sce_visible)
|
||||
with gr.Box():
|
||||
gr.Markdown(self.component_names.resolution_param)
|
||||
with gr.Group():
|
||||
gr.Markdown(self.component_names.resolution_param)
|
||||
with gr.Row(equal_height=True):
|
||||
self.resolution_height = gr.Dropdown(
|
||||
choices=list(self.h_level_dict.keys()),
|
||||
value=self.para_data.get('RESOLUTION',
|
||||
1024)[0],
|
||||
label=self.component_names.resolution_height,
|
||||
allow_custom_value=True,
|
||||
interactive=True)
|
||||
self.resolution_width = gr.Dropdown(
|
||||
choices=self.h_level_dict[
|
||||
self.resolution_height.value],
|
||||
value=self.para_data.get('RESOLUTION',
|
||||
1024)[1],
|
||||
label=self.component_names.resolution_width,
|
||||
allow_custom_value=True,
|
||||
interactive=True)
|
||||
self.enable_resolution_bucket = gr.Checkbox(
|
||||
value=False,
|
||||
container=True,
|
||||
interactive=True,
|
||||
label=self.component_names.
|
||||
enable_resolution_bucket,
|
||||
info=self.component_names.
|
||||
enable_resolution_bucket_ins)
|
||||
with gr.Column(visible=self.enable_resolution_bucket.
|
||||
value) as self.resolution_bucket_param:
|
||||
with gr.Row():
|
||||
self.min_bucket_resolution = gr.Number(
|
||||
label=self.component_names.
|
||||
min_bucket_resolution,
|
||||
value=self.para_data.get(
|
||||
'min_bucket_resolution', 256),
|
||||
interactive=True)
|
||||
self.max_bucket_resolution = gr.Number(
|
||||
label=self.component_names.
|
||||
max_bucket_resolution,
|
||||
value=self.para_data.get(
|
||||
'max_bucket_resolution', 1024),
|
||||
interactive=True)
|
||||
with gr.Row(equal_height=True):
|
||||
self.resolution_height = gr.Dropdown(
|
||||
choices=list(self.h_level_dict.keys()),
|
||||
value=self.para_data.get(
|
||||
'RESOLUTION', 1024)[0],
|
||||
self.bucket_resolution_steps = gr.Number(
|
||||
label=self.component_names.
|
||||
resolution_height,
|
||||
allow_custom_value=True,
|
||||
interactive=True)
|
||||
self.resolution_width = gr.Dropdown(
|
||||
choices=self.h_level_dict[
|
||||
self.resolution_height.value],
|
||||
bucket_resolution_steps,
|
||||
value=self.para_data.get(
|
||||
'RESOLUTION', 1024)[1],
|
||||
label=self.component_names.
|
||||
resolution_width,
|
||||
allow_custom_value=True,
|
||||
'bucket_resolution_steps', 64),
|
||||
interactive=True)
|
||||
with gr.Box():
|
||||
self.enable_resolution_bucket = gr.Checkbox(
|
||||
with gr.Group():
|
||||
self.bucket_no_upscale = gr.Checkbox(
|
||||
value=False,
|
||||
container=True,
|
||||
interactive=True,
|
||||
label=self.component_names.enable_resolution_bucket,
|
||||
info=self.component_names.enable_resolution_bucket_ins)
|
||||
with gr.Column(visible=self.enable_resolution_bucket.
|
||||
value) as self.resolution_bucket_param:
|
||||
with gr.Row():
|
||||
self.min_bucket_resolution = gr.Number(
|
||||
label=self.component_names.
|
||||
min_bucket_resolution,
|
||||
value=self.para_data.get(
|
||||
'min_bucket_resolution', 256),
|
||||
interactive=True)
|
||||
self.max_bucket_resolution = gr.Number(
|
||||
label=self.component_names.
|
||||
max_bucket_resolution,
|
||||
value=self.para_data.get(
|
||||
'max_bucket_resolution', 1024),
|
||||
interactive=True)
|
||||
with gr.Row(equal_height=True):
|
||||
self.bucket_resolution_steps = gr.Number(
|
||||
label=self.component_names.
|
||||
bucket_resolution_steps,
|
||||
value=self.para_data.get(
|
||||
'bucket_resolution_steps', 64),
|
||||
interactive=True)
|
||||
with gr.Box():
|
||||
self.bucket_no_upscale = gr.Checkbox(
|
||||
value=False,
|
||||
container=True,
|
||||
interactive=True,
|
||||
label=self.component_names.bucket_no_upscale,
|
||||
info=self.component_names.bucket_no_upscale_ins)
|
||||
|
||||
bucket_no_upscale,
|
||||
info=self.component_names.
|
||||
bucket_no_upscale_ins)
|
||||
with gr.Group():
|
||||
gr.Markdown(self.component_names.base_param)
|
||||
with gr.Row():
|
||||
self.train_epoch = gr.Number(
|
||||
label=self.component_names.train_epoch,
|
||||
@@ -303,13 +300,11 @@ class TrainerUI(UIBase):
|
||||
with gr.Row():
|
||||
self.save_interval = gr.Number(
|
||||
label=self.component_names.save_interval,
|
||||
value=self.para_data.get(
|
||||
'SAVE_INTERVAL', 10),
|
||||
value=self.para_data.get('SAVE_INTERVAL', 10),
|
||||
precision=0,
|
||||
interactive=True)
|
||||
self.train_batch_size = gr.Number(
|
||||
label=self.component_names.
|
||||
train_batch_size,
|
||||
label=self.component_names.train_batch_size,
|
||||
value=self.para_data.get(
|
||||
'TRAIN_BATCH_SIZE', 4),
|
||||
precision=0,
|
||||
@@ -318,14 +313,12 @@ class TrainerUI(UIBase):
|
||||
with gr.Row():
|
||||
self.prompt_prefix = gr.Text(
|
||||
label=self.component_names.prompt_prefix,
|
||||
value=self.para_data.get(
|
||||
'TRAIN_PREFIX', ''))
|
||||
value=self.para_data.get('TRAIN_PREFIX', ''))
|
||||
self.replace_keywords = gr.Text(
|
||||
label=self.component_names.
|
||||
replace_keywords,
|
||||
label=self.component_names.replace_keywords,
|
||||
value='')
|
||||
with gr.Row(variant='panel', equal_height=True):
|
||||
with gr.Box():
|
||||
with gr.Group():
|
||||
with gr.Row(variant='panel', equal_height=True):
|
||||
gr.Markdown(self.component_names.work_name)
|
||||
with gr.Row(variant='panel', equal_height=True):
|
||||
@@ -386,9 +379,9 @@ class TrainerUI(UIBase):
|
||||
self.component_names.data_source_choices[0],
|
||||
self.component_names.data_source_choices[2]
|
||||
]:
|
||||
return gr.Box(visible=False)
|
||||
return gr.Group(visible=False)
|
||||
elif data_source == self.component_names.data_source_choices[1]:
|
||||
return gr.Box(visible=True)
|
||||
return gr.Group(visible=True)
|
||||
|
||||
self.data_source.change(fn=change_data_source,
|
||||
inputs=[self.data_source],
|
||||
@@ -826,8 +819,10 @@ class TrainerUI(UIBase):
|
||||
cfg = current_val
|
||||
# update config
|
||||
cfg['SOLVER']['WORK_DIR'] = work_dir
|
||||
# cfg['SOLVER']['OPTIMIZER']['LEARNING_RATE'] = float(
|
||||
# learning_rate * 640 / int(train_batch_size))
|
||||
cfg['SOLVER']['OPTIMIZER']['LEARNING_RATE'] = float(
|
||||
learning_rate * 640 / int(train_batch_size))
|
||||
learning_rate)
|
||||
cfg['SOLVER']['MAX_EPOCHS'] = int(train_epoch)
|
||||
cfg['SOLVER']['TRAIN_DATA']['BATCH_SIZE'] = int(
|
||||
train_batch_size)
|
||||
@@ -903,8 +898,8 @@ class TrainerUI(UIBase):
|
||||
if work_name not in inference_ui.model_list:
|
||||
inference_ui.model_list.append(work_name)
|
||||
gr.Info('Start Training!' + message)
|
||||
return gr.Dropdown.update(choices=inference_ui.model_list,
|
||||
value=work_name)
|
||||
return gr.Dropdown(choices=inference_ui.model_list,
|
||||
value=work_name)
|
||||
|
||||
self.training_button.click(
|
||||
run_train,
|
||||
|
||||
@@ -5,9 +5,8 @@ import shutil
|
||||
from collections import OrderedDict
|
||||
|
||||
import gradio as gr
|
||||
from huggingface_hub import HfApi, snapshot_download
|
||||
|
||||
import scepter
|
||||
from huggingface_hub import HfApi, snapshot_download
|
||||
from scepter.modules.utils.config import Config
|
||||
from scepter.modules.utils.file_system import FS
|
||||
from scepter.modules.utils.module_transform import \
|
||||
@@ -95,7 +94,7 @@ class BrowserUI(UIBase):
|
||||
diffusion_models_choice, diffusion_model, tuner_models_choice, tuner_model = self.get_choices_and_values(
|
||||
)
|
||||
with gr.Column():
|
||||
with gr.Box():
|
||||
with gr.Group():
|
||||
gr.Markdown(self.component_names.browser_block_name)
|
||||
with gr.Row(variant='panel', equal_height=True):
|
||||
with gr.Column(scale=4, min_width=0):
|
||||
@@ -114,38 +113,33 @@ class BrowserUI(UIBase):
|
||||
interactive=True)
|
||||
with gr.Column(scale=1, min_width=0):
|
||||
self.save_button = gr.Button(
|
||||
label='Save',
|
||||
value=self.component_names.save_symbol,
|
||||
elem_classes='type_row',
|
||||
elem_id='save_button',
|
||||
visible=True)
|
||||
self.delete_button = gr.Button(
|
||||
label='Delete',
|
||||
value=self.component_names.delete_symbol,
|
||||
elem_classes='type_row',
|
||||
elem_id='delete_button',
|
||||
visible=False)
|
||||
self.model_export = gr.Button(
|
||||
label='Model Export',
|
||||
value=self.component_names.model_export,
|
||||
elem_classes='type_row',
|
||||
elem_id='save_button',
|
||||
visible=True)
|
||||
with gr.Column(scale=1, min_width=0):
|
||||
self.refresh_button = gr.Button(
|
||||
label='Refresh',
|
||||
value=self.component_names.refresh_symbol,
|
||||
elem_classes='type_row',
|
||||
elem_id='refresh_button',
|
||||
visible=True)
|
||||
self.model_import = gr.Button(
|
||||
label='Model Import',
|
||||
value=self.component_names.model_import,
|
||||
elem_classes='type_row',
|
||||
elem_id='save_button',
|
||||
visible=True)
|
||||
|
||||
with gr.Box(visible=False) as self.export_setting:
|
||||
with gr.Group(visible=False) as self.export_setting:
|
||||
gr.Markdown(self.component_names.export_desc)
|
||||
with gr.Column(variant='panel'):
|
||||
with gr.Row(equal_height=True):
|
||||
@@ -157,7 +151,6 @@ class BrowserUI(UIBase):
|
||||
show_label=False)
|
||||
with gr.Column(scale=1, min_width=0):
|
||||
self.export_close = gr.Button(
|
||||
label='Close Export',
|
||||
value=self.component_names.close,
|
||||
elem_classes='type_row',
|
||||
elem_id='save_button')
|
||||
@@ -185,7 +178,6 @@ class BrowserUI(UIBase):
|
||||
value=False)
|
||||
with gr.Column(scale=1, min_width=0):
|
||||
self.ms_export_submit = gr.Button(
|
||||
label='Submit MS',
|
||||
value=self.component_names.submit,
|
||||
elem_classes='type_row',
|
||||
elem_id='save_button')
|
||||
@@ -213,12 +205,11 @@ class BrowserUI(UIBase):
|
||||
value=False)
|
||||
with gr.Column(scale=1, min_width=0):
|
||||
self.hf_export_submit = gr.Button(
|
||||
label='Submit HF',
|
||||
value=self.component_names.submit,
|
||||
elem_classes='type_row',
|
||||
elem_id='save_button')
|
||||
|
||||
with gr.Box(visible=False) as self.import_setting:
|
||||
with gr.Group(visible=False) as self.import_setting:
|
||||
gr.Markdown(self.component_names.import_desc)
|
||||
with gr.Column(variant='panel'):
|
||||
with gr.Row(equal_height=True):
|
||||
@@ -232,21 +223,20 @@ class BrowserUI(UIBase):
|
||||
show_label=False)
|
||||
with gr.Column(scale=1, min_width=0):
|
||||
self.import_close = gr.Button(
|
||||
label='Close Download',
|
||||
value=self.component_names.close,
|
||||
elem_classes='type_row',
|
||||
elem_id='save_button')
|
||||
|
||||
with gr.Row(
|
||||
equal_height=True) as self.ms_import_setting:
|
||||
with gr.Column(scale=4.5, min_width=0):
|
||||
with gr.Column(scale=4, min_width=0):
|
||||
self.ms_modelid = gr.Text(
|
||||
label=self.component_names.ms_modelid,
|
||||
show_label=False,
|
||||
container=False,
|
||||
placeholder='ModelScope Model Path',
|
||||
value='')
|
||||
with gr.Column(scale=4.5, min_width=0):
|
||||
with gr.Column(scale=4, min_width=0):
|
||||
self.ms_import_username = gr.Text(
|
||||
label=self.component_names.ms_username,
|
||||
show_label=False,
|
||||
@@ -255,7 +245,6 @@ class BrowserUI(UIBase):
|
||||
value='')
|
||||
with gr.Column(scale=1, min_width=0):
|
||||
self.ms_import_submit = gr.Button(
|
||||
label='Submit MS',
|
||||
value=self.component_names.submit,
|
||||
elem_classes='type_row',
|
||||
elem_id='save_button')
|
||||
@@ -285,7 +274,6 @@ class BrowserUI(UIBase):
|
||||
value='')
|
||||
with gr.Column(scale=1, min_width=0):
|
||||
self.hf_import_submit = gr.Button(
|
||||
label='Submit HF',
|
||||
value=self.component_names.submit,
|
||||
elem_classes='type_row',
|
||||
elem_id='save_button')
|
||||
@@ -333,7 +321,6 @@ class BrowserUI(UIBase):
|
||||
placeholder='Upload Tuner Name',
|
||||
value='')
|
||||
self.local_upload_bt = gr.Button(
|
||||
label='Submit Local Model',
|
||||
value=self.component_names.submit,
|
||||
elem_classes='type_row',
|
||||
elem_id='upload_button')
|
||||
@@ -358,7 +345,8 @@ class BrowserUI(UIBase):
|
||||
tar_path = os.path.join(tar_path, tuner_name)
|
||||
if FS.exists(tar_path):
|
||||
raise gr.Error(self.component_names.same_name)
|
||||
local_model_dir, _ = FS.map_to_local(tar_path)
|
||||
# local_model_dir, _ = FS.map_to_local(tar_path)
|
||||
local_model_dir = src_path
|
||||
FS.put_dir_from_local_dir(src_path, tar_path)
|
||||
# save image
|
||||
tuner_example_path = None
|
||||
|
||||
@@ -43,7 +43,7 @@ class InfoUI(UIBase):
|
||||
|
||||
def create_ui(self, *args, **kwargs):
|
||||
with gr.Column():
|
||||
with gr.Box():
|
||||
with gr.Group():
|
||||
gr.Markdown(self.component_names.info_block_name)
|
||||
with gr.Row(variant='panel', equal_height=True):
|
||||
with gr.Column(variant='panel', scale=1, min_width=0):
|
||||
@@ -80,7 +80,6 @@ class InfoUI(UIBase):
|
||||
with gr.Column(scale=1):
|
||||
self.save_bt2 = gr.Button(
|
||||
value=self.component_names.save_symbol,
|
||||
label='Save',
|
||||
elem_classes='type_row',
|
||||
elem_id='save_button',
|
||||
visible=True,
|
||||
@@ -91,7 +90,7 @@ class InfoUI(UIBase):
|
||||
with gr.Row(equal_height=True):
|
||||
self.tuner_example = gr.Image(
|
||||
label=self.component_names.tuner_example,
|
||||
source='upload',
|
||||
sources=['upload'],
|
||||
value=None,
|
||||
interactive=True)
|
||||
with gr.Row(equal_height=True):
|
||||
@@ -116,7 +115,6 @@ class InfoUI(UIBase):
|
||||
self.component_names.go_to_inference)
|
||||
with gr.Column(scale=1, min_width=0):
|
||||
self.local_download_bt = gr.Button(
|
||||
label='Download to Local Dir',
|
||||
value=self.component_names.
|
||||
download_to_local,
|
||||
# elem_classes='type_row',
|
||||
@@ -134,7 +132,6 @@ class InfoUI(UIBase):
|
||||
show_label=False)
|
||||
with gr.Column(scale=1, min_width=0):
|
||||
self.download_confirm = gr.Button(
|
||||
label='Confirm Download Format',
|
||||
value=self.component_names.submit,
|
||||
elem_classes='type_row',
|
||||
elem_id='save_button')
|
||||
|
||||
@@ -2,8 +2,6 @@
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import argparse
|
||||
import importlib
|
||||
# from safetensors.torch import save_file
|
||||
# from safetensors import safe_open
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
|
||||
@@ -16,7 +16,7 @@ parser.add_argument(
|
||||
'--python_engine',
|
||||
dest='python_engine',
|
||||
help='the engine path of python interpreter!',
|
||||
default='',
|
||||
default='python',
|
||||
type=str,
|
||||
)
|
||||
|
||||
@@ -24,7 +24,7 @@ parser.add_argument(
|
||||
'--script',
|
||||
dest='script',
|
||||
help='the script to run!',
|
||||
default='',
|
||||
default='main_mmpose.py',
|
||||
type=str,
|
||||
)
|
||||
parser.add_argument(
|
||||
|
||||
@@ -171,5 +171,4 @@ if __name__ == '__main__':
|
||||
root_path=config['ROOT'],
|
||||
show_error=True,
|
||||
debug=True,
|
||||
enable_queue=True,
|
||||
auth=check_auth if len(auth_info) > 0 else None)
|
||||
|
||||
+1
-1
@@ -1,7 +1,7 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
|
||||
__version__ = '1.0.2'
|
||||
__version__ = '1.0.3'
|
||||
|
||||
version_info = tuple(int(x) for x in __version__.split('.')[0:3])
|
||||
|
||||
|
||||
Reference in New Issue
Block a user