diff --git a/scepter/__init__.py b/scepter/__init__.py index 3613881..314d619 100644 --- a/scepter/__init__.py +++ b/scepter/__init__.py @@ -1,18 +1,28 @@ # -*- coding: utf-8 -*- # Copyright (c) Alibaba, Inc. and its affiliates. -import os +from typing import TYPE_CHECKING +from scepter.modules.utils.import_utils import LazyImportModule -import scepter -from scepter.modules import data, model, opt, solver, transform, utils -from scepter.tools.helper import get_module_list as module_list -from scepter.tools.helper import \ - get_module_object_config as configures_by_objects -from scepter.tools.helper import get_module_objects as objects_by_module -from scepter.version import __version__, version_info -dirname = os.path.dirname(scepter.__file__) +if TYPE_CHECKING: + from scepter.modules import data, model, opt, solver, transform, utils + from scepter.tools.helper import get_module_list as module_list + from scepter.tools.helper import \ + get_module_object_config as configures_by_objects + from scepter.tools.helper import get_module_objects as objects_by_module + from scepter.version import __version__, version_info +else: + _import_structure = { + 'modules': ['data', 'model', 'opt', 'solver', 'transform', 'utils'], + 'helper': ['get_module_list', 'get_module_object_config', 'get_module_objects'], + 'version': ['__version__', 'version_info'] + } -__all__ = [ - utils, transform, data, model, solver, version_info, opt, '__version__', - 'dirname' -] + import sys + sys.modules[__name__] = LazyImportModule( + __name__, + globals()['__file__'], + _import_structure, + module_spec=__spec__, + extra_objects={}, + ) diff --git a/scepter/methods/examples/generation/dit_cogvideox1.5_5b_i2v_lora.yaml b/scepter/methods/examples/generation/dit_cogvideox1.5_5b_i2v_lora.yaml new file mode 100644 index 0000000..03b5696 --- /dev/null +++ b/scepter/methods/examples/generation/dit_cogvideox1.5_5b_i2v_lora.yaml @@ -0,0 +1,277 @@ +ENV: + BACKEND: nccl + SEED: 42 + TENSOR_PARALLEL_SIZE: 1 + PIPELINE_PARALLEL_SIZE: 1 + SYS_ENVS: + TORCH_CUDNN_V8_API_ENABLED: '1' + TOKENIZERS_PARALLELISM: 'false' + TF_CPP_MIN_LOG_LEVEL: '3' + PYTORCH_CUDA_ALLOC_CONF: 'expandable_segments:True' +# +SOLVER: + NAME: LatentDiffusionVideoSolver + MAX_STEPS: 2000 + USE_AMP: True + DTYPE: bfloat16 + USE_FAIRSCALE: False + USE_FSDP: True + LOAD_MODEL_ONLY: False + ENABLE_GRADSCALER: False + USE_SCALER: False + RESUME_FROM: + WORK_DIR: ./cache/save_data/dit_cogvideox1.5_5b_i2v_lora + LOG_FILE: std_log.txt + EVAL_INTERVAL: 100 + LOG_TRAIN_NUM: 4 + FPS: 16 + SHARDING_STRATEGY: full_shard + FSDP_REDUCE_DTYPE: float32 + FSDP_BUFFER_DTYPE: float32 + FSDP_SHARD_MODULES: [ 'model', 'cond_stage_model.model'] + SAVE_MODULES: [ 'model', 'cond_stage_model.model'] + TRAIN_MODULES: ['model'] + # + FILE_SYSTEM: + NAME: "ModelscopeFs" + TEMP_DIR: "./cache/cache_data" + # + TUNER: + - NAME: SwiftLoRA + R: 64 + LORA_ALPHA: 64 + LORA_DROPOUT: 0.0 + BIAS: "none" + TARGET_MODULES: "model.*(.to_k|.to_q|.to_v|.to_out.0)$" + # + MODEL: + NAME: LatentDiffusionCogVideoX + PRETRAINED_MODEL: + PARAMETERIZATION: v + TIMESTEPS: 1000 + MIN_SNR_GAMMA: 3.0 + ZERO_TERMINAL_SNR: True + SCALE_FACTOR_SPATIAL: 8 + SCALE_FACTOR_TEMPORAL: 4 + SCALING_FACTOR_IMAGE: 0.7 + NOISED_IMAGE_DROPOUT: 0.05 + INVERT_SCALE_LATENTS: True + IGNORE_KEYS: [ ] + DEFAULT_N_PROMPT: + USE_EMA: False + EVAL_EMA: False + DIFFUSION: + NAME: BaseDiffusion + PREDICTION_TYPE: v + USE_DYNAMIC_CFG: False + NOISE_SCHEDULER: + NAME: ScaledLinearScheduler + BETA_MIN: 0.00085 + BETA_MAX: 0.012 + SNR_SHIFT_SCALE: 1.0 + RESCALE_BETAS_ZERO_SNR: True + DIFFUSION_SAMPLERS: + NAME: DDIMSampler + DISCRETIZATION_TYPE: trailing + ETA: 0.0 + # + DIFFUSION_MODEL: + NAME: CogVideoXTransformer3DModel + DTYPE: bfloat16 + PRETRAINED_MODEL: # 5b-I2V diff + - ms://ZhipuAI/CogVideoX1.5-5B-I2V@transformer/diffusion_pytorch_model-00001-of-00003.safetensors + - ms://ZhipuAI/CogVideoX1.5-5B-I2V@transformer/diffusion_pytorch_model-00002-of-00003.safetensors + - ms://ZhipuAI/CogVideoX1.5-5B-I2V@transformer/diffusion_pytorch_model-00003-of-00003.safetensors + NUM_ATTENTION_HEADS: 48 + ATTENTION_HEAD_DIM: 64 + IN_CHANNELS: 32 + LATENT_CHANNELS: 16 + OUT_CHANNELS: 16 + FLIP_SIN_TO_COS: True + FREQ_SHIFT: 0 + TIME_EMBED_DIM: 512 + TEXT_EMBED_DIM: 4096 + OFS_EMBED_DIM: 512 # v1.5 diff + NUM_LAYERS: 42 + DROPOUT: 0.0 + ATTENTION_BIAS: True + SAMPLE_WIDTH: 300 + SAMPLE_HEIGHT: 300 + SAMPLE_FRAMES: 81 + PATCH_SIZE: 2 + PATCH_SIZE_T: 2 # v1.5 diff + PATCH_BIAS: False # v1.5 diff + TEMPORAL_COMPRESSION_RATIO: 4 + MAX_TEXT_SEQ_LENGTH: 224 + ACTIVATION_FN: "gelu-approximate" + TIMESTEP_ACTIVATION_FN: "silu" + NORM_ELEMENTWISE_AFFINE: True + NORM_EPS: 1e-5 + SPATIAL_INTERPOLATION_SCALE: 1.875 + TEMPORAL_INTERPOLATION_SCALE: 1.0 + USE_ROTARY_POSITIONAL_EMBEDDINGS: True + USE_LEARNED_POSITIONAL_EMBEDDINGS: False + GRADIENT_CHECKPOINTING: True + # + FIRST_STAGE_MODEL: + NAME: AutoencoderKLCogVideoX + DTYPE: bfloat16 + PRETRAINED_MODEL: ms://ZhipuAI/CogVideoX1.5-5B-I2V@vae/diffusion_pytorch_model.safetensors + SAMPLE_HEIGHT: 768 + SAMPLE_WIDTH: 1360 + USE_QUANT_CONV: False + USE_POST_QUANT_CONV: False + USE_SLICING: True + USE_TILING: True + GRADIENT_CHECKPOINTING: True + ENCODER: + NAME: CogVideoXEncoder3D + IN_CHANNELS: 3 + OUT_CHANNELS: 16 + UP_BLOCK_TYPES: [ "CogVideoXDownBlock3D", "CogVideoXDownBlock3D", "CogVideoXDownBlock3D", "CogVideoXDownBlock3D" ] + BLOCK_OUT_CHANNELS: [ 128, 256, 256, 512 ] + LAYERS_PER_BLOCK: 3 + ACT_FN: "silu" + NORM_EPS: 1e-6 + NORM_NUM_GROUPS: 32 + DROPOUT: 0.0 + PAD_MODE: "first" + TEMPORAL_COMPRESSION_RATIO: 4 + GRADIENT_CHECKPOINTING: True + DECODER: + NAME: CogVideoXDecoder3D + IN_CHANNELS: 16 + OUT_CHANNELS: 3 + UP_BLOCK_TYPES: [ "CogVideoXUpBlock3D", "CogVideoXUpBlock3D", "CogVideoXUpBlock3D", "CogVideoXUpBlock3D" ] + BLOCK_OUT_CHANNELS: [ 128, 256, 256, 512 ] + LAYERS_PER_BLOCK: 3 + ACT_FN: "silu" + NORM_EPS: 1e-6 + NORM_NUM_GROUPS: 32 + DROPOUT: 0.0 + PAD_MODE: "first" + TEMPORAL_COMPRESSION_RATIO: 4 + GRADIENT_CHECKPOINTING: True + # + COND_STAGE_MODEL: + NAME: T5EmbedderHF + PRETRAINED_MODEL: ms://AI-ModelScope/t5-v1_1-xxl + TOKENIZER_PATH: ms://AI-ModelScope/t5-v1_1-xxl + LENGTH: 224 + CLEAN: + USE_GRAD: False + T5_DTYPE: bfloat16 + # + LOSS: + NAME: ReconstructLoss + LOSS_TYPE: l2 + # + SAMPLE_ARGS: + SAMPLER: ddim + SAMPLE_STEPS: 50 + SEED: 42 + GUIDE_SCALE: 6.0 + GUIDE_RESCALE: 0.0 + NUM_FRAMES: 81 + IMAGE_SIZE: [768, 1360] + # + OPTIMIZER: + NAME: Adam + LEARNING_RATE: 1e-3 + BETAS: [ 0.9, 0.95 ] + EPS: 1e-8 + WEIGHT_DECAY: 0.0 + AMSGRAD: False + # +# LR_SCHEDULER: +# NAME: StepAnnealingLR +# WARMUP_STEPS: 200 +# TOTAL_STEPS: 2000 +# DECAY_MODE: 'cosine' + # + TRAIN_DATA: + NAME: VideoGenDataset + MODE: train + PIN_MEMORY: True + BATCH_SIZE: 1 + NUM_WORKERS: 0 + NUM_FRAMES: 85 + FPS: 16 + HEIGHT: 768 + WIDTH: 1360 + PROMPT_PREFIX: 'DISNEY ' + DATA_TYPE: 'i2v' + SAMPLER: + NAME: MixtureOfSamplers + SUB_SAMPLERS: + - NAME: MultiLevelBatchSampler + PROB: 1.0 + FIELDS: [ "video_path", "prompt" ] + DELIMITER: '#;#' + PATH_PREFIX: cache/datasets/Disney-VideoGeneration-Dataset/ + INDEX_FILE: cache/datasets/Disney-VideoGeneration-Dataset/index.jsonl + TRANSFORMS: + - NAME: Select + KEYS: [ "video", "image", "prompt" ] + META_KEYS: [ ] + # +# EVAL_DATA: +# NAME: Text2ImageDataset +# MODE: eval +# PROMPT_FILE: +# PROMPT_DATA: [ "A cat running.#;#asset/images/edit_tuner/cat_512.jpg" ] +# FIELDS: [ "prompt", "img_path" ] +# DELIMITER: '#;#' +# PROMPT_PREFIX: '' +# PIN_MEMORY: True +# BATCH_SIZE: 1 +# USE_NUM: 8 +# NUM_WORKERS: 0 +# IMAGE_SIZE: [768, 1360] +# TRANSFORMS: +# - NAME: LoadImageFromFileList +# FILE_KEYS: [ 'img_path' ] +# RGB_ORDER: RGB +# BACKEND: pillow +# - NAME: FlexibleResize +# INTERPOLATION: bilinear +# SIZE: [768, 1360] +# INPUT_KEY: [ 'img' ] +# OUTPUT_KEY: [ 'img' ] +# BACKEND: pillow +# - NAME: FlexibleCenterCrop +# SIZE: [768, 1360] +# 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: [ 'image_size' ] + # + TRAIN_HOOKS: + - NAME: ProbeDataHook + PROB_INTERVAL: 100 + PRIORITY: 0 + - NAME: BackwardHook + PRIORITY: 10 + - NAME: LogHook + LOG_INTERVAL: 10 + PRIORITY: 20 + - NAME: CheckpointHook + INTERVAL: 1000 + PRIORITY: 40 + # +# EVAL_HOOKS: +# - NAME: ProbeDataHook +# PROB_INTERVAL: 100 +# PRIORITY: 0 \ No newline at end of file diff --git a/scepter/methods/examples/generation/dit_cogvideox1.5_5b_lora.yaml b/scepter/methods/examples/generation/dit_cogvideox1.5_5b_lora.yaml new file mode 100644 index 0000000..dce5102 --- /dev/null +++ b/scepter/methods/examples/generation/dit_cogvideox1.5_5b_lora.yaml @@ -0,0 +1,248 @@ +ENV: + BACKEND: nccl + SEED: 42 + TENSOR_PARALLEL_SIZE: 1 + PIPELINE_PARALLEL_SIZE: 1 + SYS_ENVS: + TORCH_CUDNN_V8_API_ENABLED: '1' + TOKENIZERS_PARALLELISM: 'false' + TF_CPP_MIN_LOG_LEVEL: '3' + PYTORCH_CUDA_ALLOC_CONF: 'expandable_segments:True' +# +SOLVER: + NAME: LatentDiffusionVideoSolver + MAX_STEPS: 2000 + USE_AMP: True + DTYPE: bfloat16 + USE_FAIRSCALE: False + USE_FSDP: True + LOAD_MODEL_ONLY: False + ENABLE_GRADSCALER: False + USE_SCALER: False + RESUME_FROM: + WORK_DIR: ./cache/save_data/dit_cogvideox1.5_5b_lora + LOG_FILE: std_log.txt + EVAL_INTERVAL: 100 + LOG_TRAIN_NUM: 4 + FPS: 16 + SHARDING_STRATEGY: full_shard + FSDP_REDUCE_DTYPE: float32 + FSDP_BUFFER_DTYPE: float32 + FSDP_SHARD_MODULES: [ 'model', 'cond_stage_model.model'] + SAVE_MODULES: [ 'model', 'cond_stage_model.model'] + TRAIN_MODULES: ['model'] + # + FILE_SYSTEM: + NAME: "ModelscopeFs" + TEMP_DIR: "./cache/cache_data" + # + TUNER: + - NAME: SwiftLoRA + R: 64 + LORA_ALPHA: 64 + LORA_DROPOUT: 0.0 + BIAS: "none" + TARGET_MODULES: "model.*(.to_k|.to_q|.to_v|.to_out.0)$" + # + MODEL: + NAME: LatentDiffusionCogVideoX + PRETRAINED_MODEL: + PARAMETERIZATION: v + TIMESTEPS: 1000 + MIN_SNR_GAMMA: 3.0 + ZERO_TERMINAL_SNR: True + SCALE_FACTOR_SPATIAL: 8 + SCALE_FACTOR_TEMPORAL: 4 + SCALING_FACTOR_IMAGE: 0.7 + INVERT_SCALE_LATENTS: True + IGNORE_KEYS: [ ] + DEFAULT_N_PROMPT: + USE_EMA: False + EVAL_EMA: False + DIFFUSION: + NAME: BaseDiffusion + PREDICTION_TYPE: v + USE_DYNAMIC_CFG: False + NOISE_SCHEDULER: + NAME: ScaledLinearScheduler + BETA_MIN: 0.00085 + BETA_MAX: 0.012 + SNR_SHIFT_SCALE: 1.0 + RESCALE_BETAS_ZERO_SNR: True + DIFFUSION_SAMPLERS: + NAME: DDIMSampler + DISCRETIZATION_TYPE: trailing + ETA: 0.0 + # + DIFFUSION_MODEL: + NAME: CogVideoXTransformer3DModel + DTYPE: bfloat16 + PRETRAINED_MODEL: + - ms://ZhipuAI/CogVideoX1.5-5B@transformer/diffusion_pytorch_model-00001-of-00003.safetensors + - ms://ZhipuAI/CogVideoX1.5-5B@transformer/diffusion_pytorch_model-00002-of-00003.safetensors + - ms://ZhipuAI/CogVideoX1.5-5B@transformer/diffusion_pytorch_model-00003-of-00003.safetensors + NUM_ATTENTION_HEADS: 48 + ATTENTION_HEAD_DIM: 64 + IN_CHANNELS: 16 + OUT_CHANNELS: 16 + FLIP_SIN_TO_COS: True + FREQ_SHIFT: 0 + TIME_EMBED_DIM: 512 + TEXT_EMBED_DIM: 4096 + NUM_LAYERS: 42 + DROPOUT: 0.0 + ATTENTION_BIAS: True + SAMPLE_WIDTH: 300 + SAMPLE_HEIGHT: 300 + SAMPLE_FRAMES: 81 + PATCH_SIZE: 2 + PATCH_SIZE_T: 2 # v1.5 diff + PATCH_BIAS: False # v1.5 diff + TEMPORAL_COMPRESSION_RATIO: 4 + MAX_TEXT_SEQ_LENGTH: 224 + ACTIVATION_FN: "gelu-approximate" + TIMESTEP_ACTIVATION_FN: "silu" + NORM_ELEMENTWISE_AFFINE: True + NORM_EPS: 1e-5 + SPATIAL_INTERPOLATION_SCALE: 1.875 + TEMPORAL_INTERPOLATION_SCALE: 1.0 + USE_ROTARY_POSITIONAL_EMBEDDINGS: True + USE_LEARNED_POSITIONAL_EMBEDDINGS: False + GRADIENT_CHECKPOINTING: True + # + FIRST_STAGE_MODEL: + NAME: AutoencoderKLCogVideoX + DTYPE: bfloat16 + PRETRAINED_MODEL: ms://ZhipuAI/CogVideoX1.5-5B@vae/diffusion_pytorch_model.safetensors + SAMPLE_HEIGHT: 768 + SAMPLE_WIDTH: 1360 + USE_QUANT_CONV: False + USE_POST_QUANT_CONV: False + USE_SLICING: True + USE_TILING: True + GRADIENT_CHECKPOINTING: True + ENCODER: + NAME: CogVideoXEncoder3D + IN_CHANNELS: 3 + OUT_CHANNELS: 16 + UP_BLOCK_TYPES: [ "CogVideoXDownBlock3D", "CogVideoXDownBlock3D", "CogVideoXDownBlock3D", "CogVideoXDownBlock3D" ] + BLOCK_OUT_CHANNELS: [ 128, 256, 256, 512 ] + LAYERS_PER_BLOCK: 3 + ACT_FN: "silu" + NORM_EPS: 1e-6 + NORM_NUM_GROUPS: 32 + DROPOUT: 0.0 + PAD_MODE: "first" + TEMPORAL_COMPRESSION_RATIO: 4 + GRADIENT_CHECKPOINTING: True + DECODER: + NAME: CogVideoXDecoder3D + IN_CHANNELS: 16 + OUT_CHANNELS: 3 + UP_BLOCK_TYPES: [ "CogVideoXUpBlock3D", "CogVideoXUpBlock3D", "CogVideoXUpBlock3D", "CogVideoXUpBlock3D" ] + BLOCK_OUT_CHANNELS: [ 128, 256, 256, 512 ] + LAYERS_PER_BLOCK: 3 + ACT_FN: "silu" + NORM_EPS: 1e-6 + NORM_NUM_GROUPS: 32 + DROPOUT: 0.0 + PAD_MODE: "first" + TEMPORAL_COMPRESSION_RATIO: 4 + GRADIENT_CHECKPOINTING: True + # + COND_STAGE_MODEL: + NAME: T5EmbedderHF + PRETRAINED_MODEL: ms://AI-ModelScope/t5-v1_1-xxl + TOKENIZER_PATH: ms://AI-ModelScope/t5-v1_1-xxl + LENGTH: 224 + CLEAN: + USE_GRAD: False + T5_DTYPE: bfloat16 + # + LOSS: + NAME: ReconstructLoss + LOSS_TYPE: l2 + # + SAMPLE_ARGS: + SAMPLER: ddim + SAMPLE_STEPS: 50 + SEED: 42 + GUIDE_SCALE: 6.0 + GUIDE_RESCALE: 0.0 + NUM_FRAMES: 81 + IMAGE_SIZE: [768, 1360] + # + OPTIMIZER: + NAME: Adam + LEARNING_RATE: 1e-3 + BETAS: [ 0.9, 0.95 ] + EPS: 1e-8 + WEIGHT_DECAY: 0.0 + AMSGRAD: False + # +# LR_SCHEDULER: +# NAME: StepAnnealingLR +# WARMUP_STEPS: 200 +# TOTAL_STEPS: 2000 +# DECAY_MODE: 'cosine' + # + TRAIN_DATA: + NAME: VideoGenDataset + MODE: train + PIN_MEMORY: True + BATCH_SIZE: 1 + NUM_WORKERS: 0 + NUM_FRAMES: 85 + FPS: 16 + HEIGHT: 768 + WIDTH: 1360 + PROMPT_PREFIX: 'DISNEY ' + SAMPLER: + NAME: MixtureOfSamplers + SUB_SAMPLERS: + - NAME: MultiLevelBatchSampler + PROB: 1.0 + FIELDS: [ "video_path", "prompt" ] + DELIMITER: '#;#' + PATH_PREFIX: cache/datasets/Disney-VideoGeneration-Dataset/ + INDEX_FILE: cache/datasets/Disney-VideoGeneration-Dataset/index.jsonl + TRANSFORMS: + - NAME: Select + KEYS: [ 'video', "prompt" ] + META_KEYS: [ ] + # + EVAL_DATA: + NAME: Text2ImageDataset + MODE: eval + PROMPT_FILE: + PROMPT_DATA: [ "A girl riding a bike." ] + IMAGE_SIZE: [ 768, 1360 ] + FIELDS: [ "prompt" ] + DELIMITER: '#;#' + PROMPT_PREFIX: 'DISNEY ' # '' + PIN_MEMORY: True + BATCH_SIZE: 1 + USE_NUM: 8 + NUM_WORKERS: 0 + TRANSFORMS: + - NAME: Select + KEYS: [ 'index', 'prompt' ] + META_KEYS: [ 'image_size' ] + # + TRAIN_HOOKS: + - NAME: ProbeDataHook + PROB_INTERVAL: 100 + PRIORITY: 0 + - NAME: BackwardHook + PRIORITY: 10 + - NAME: LogHook + LOG_INTERVAL: 10 + PRIORITY: 20 + - NAME: CheckpointHook + INTERVAL: 1000 + PRIORITY: 40 + # + EVAL_HOOKS: + - NAME: ProbeDataHook + PROB_INTERVAL: 100 + PRIORITY: 0 \ No newline at end of file diff --git a/scepter/methods/examples/generation/dit_cogvideox_2b_lora.yaml b/scepter/methods/examples/generation/dit_cogvideox_2b_lora.yaml index 3543c34..503fc82 100644 --- a/scepter/methods/examples/generation/dit_cogvideox_2b_lora.yaml +++ b/scepter/methods/examples/generation/dit_cogvideox_2b_lora.yaml @@ -183,6 +183,10 @@ SOLVER: PIN_MEMORY: True BATCH_SIZE: 1 NUM_WORKERS: 4 + NUM_FRAMES: 49 + FPS: 8 + HEIGHT: 480 + WIDTH: 720 PROMPT_PREFIX: 'DISNEY ' SAMPLER: NAME: MixtureOfSamplers diff --git a/scepter/methods/examples/generation/dit_cogvideox_5b_i2v_lora.yaml b/scepter/methods/examples/generation/dit_cogvideox_5b_i2v_lora.yaml index 3015577..5d9ee60 100644 --- a/scepter/methods/examples/generation/dit_cogvideox_5b_i2v_lora.yaml +++ b/scepter/methods/examples/generation/dit_cogvideox_5b_i2v_lora.yaml @@ -188,6 +188,10 @@ SOLVER: PIN_MEMORY: True BATCH_SIZE: 1 NUM_WORKERS: 0 + NUM_FRAMES: 49 + FPS: 8 + HEIGHT: 480 + WIDTH: 720 PROMPT_PREFIX: 'DISNEY ' DATA_TYPE: 'i2v' SAMPLER: diff --git a/scepter/methods/examples/generation/dit_cogvideox_5b_lora.yaml b/scepter/methods/examples/generation/dit_cogvideox_5b_lora.yaml index 56956f1..de6b6dd 100644 --- a/scepter/methods/examples/generation/dit_cogvideox_5b_lora.yaml +++ b/scepter/methods/examples/generation/dit_cogvideox_5b_lora.yaml @@ -185,6 +185,10 @@ SOLVER: PIN_MEMORY: True BATCH_SIZE: 1 NUM_WORKERS: 4 + NUM_FRAMES: 49 + FPS: 8 + HEIGHT: 480 + WIDTH: 720 PROMPT_PREFIX: 'DISNEY ' DELIMITER: '#;#' FIELDS: [ 'video_path', 'prompt' ] diff --git a/scepter/methods/studio/inference/dit/cogvideox_2b_pro.yaml b/scepter/methods/studio/inference/dit/cogvideox_2b_pro.yaml index 3e7c7b7..af5599d 100644 --- a/scepter/methods/studio/inference/dit/cogvideox_2b_pro.yaml +++ b/scepter/methods/studio/inference/dit/cogvideox_2b_pro.yaml @@ -148,4 +148,5 @@ MODEL: TOKENIZER_PATH: ms://AI-ModelScope/t5-v1_1-xxl LENGTH: 226 CLEAN: - USE_GRAD: False \ No newline at end of file + USE_GRAD: False + T5_DTYPE: bfloat16 \ No newline at end of file diff --git a/scepter/methods/studio/inference/dit/cogvideox_5b_pro.yaml b/scepter/methods/studio/inference/dit/cogvideox_5b_pro.yaml index d7c3a36..aaa9298 100644 --- a/scepter/methods/studio/inference/dit/cogvideox_5b_pro.yaml +++ b/scepter/methods/studio/inference/dit/cogvideox_5b_pro.yaml @@ -150,4 +150,5 @@ MODEL: TOKENIZER_PATH: ms://AI-ModelScope/t5-v1_1-xxl LENGTH: 226 CLEAN: - USE_GRAD: False \ No newline at end of file + USE_GRAD: False + T5_DTYPE: bfloat16 \ No newline at end of file diff --git a/scepter/methods/studio/inference/inference.yaml b/scepter/methods/studio/inference/inference.yaml index 4abfcdf..4011ef7 100644 --- a/scepter/methods/studio/inference/inference.yaml +++ b/scepter/methods/studio/inference/inference.yaml @@ -1,4 +1,5 @@ WORK_DIR: "inference" +SKIP_EXAMPLES: True DIFFUSION_PARAS: SAMPLE: VALUES: ['ddim', 'euler', 'euler_ancestral', 'heun', 'dpm2', diff --git a/scepter/modules/__init__.py b/scepter/modules/__init__.py index c0ccc94..5dc3317 100644 --- a/scepter/modules/__init__.py +++ b/scepter/modules/__init__.py @@ -1,4 +1,23 @@ # -*- coding: utf-8 -*- # Copyright (c) Alibaba, Inc. and its affiliates. -from scepter.modules import (data, inference, model, opt, solver, transform, - utils) +from typing import TYPE_CHECKING +from scepter.modules.utils.import_utils import LazyImportModule + + +if TYPE_CHECKING: + from scepter.modules import (data, inference, model, opt, solver, transform, + utils) +else: + _import_structure = { + 'modules': ['data', 'inference', 'model', 'opt', 'solver', + 'transform', 'utils'] + } + + import sys + sys.modules[__name__] = LazyImportModule( + __name__, + globals()['__file__'], + _import_structure, + module_spec=__spec__, + extra_objects={}, + ) diff --git a/scepter/modules/annotator/__init__.py b/scepter/modules/annotator/__init__.py index 7232daf..8262ea5 100644 --- a/scepter/modules/annotator/__init__.py +++ b/scepter/modules/annotator/__init__.py @@ -1,23 +1,60 @@ # -*- coding: utf-8 -*- # Copyright (c) Alibaba, Inc. and its affiliates. -from scepter.modules.annotator.base_annotator import GeneralAnnotator -from scepter.modules.annotator.canny import CannyAnnotator -from scepter.modules.annotator.color import ColorAnnotator -from scepter.modules.annotator.degradation import DegradationAnnotator -from scepter.modules.annotator.doodle import DoodleAnnotator -from scepter.modules.annotator.gray import GrayAnnotator -from scepter.modules.annotator.hed import HedAnnotator -from scepter.modules.annotator.identity import IdentityAnnotator -from scepter.modules.annotator.informative_drawing import ( - InfoDrawAnimeAnnotator, InfoDrawContourAnnotator, - InfoDrawOpenSketchAnnotator) -from scepter.modules.annotator.inpainting import InpaintingAnnotator -from scepter.modules.annotator.invert import InvertAnnotator -from scepter.modules.annotator.midas_op import MidasDetector -from scepter.modules.annotator.mlsd_op import MLSDdetector -from scepter.modules.annotator.openpose import OpenposeAnnotator -from scepter.modules.annotator.outpainting import OutpaintingAnnotator, OutpaintingResize -from scepter.modules.annotator.pidinet import PiDiAnnotator -from scepter.modules.annotator.segmentation import ESAMAnnotator -from scepter.modules.annotator.sketch import SketchAnnotator -from scepter.modules.annotator.lama import LamaAnnotator +from typing import TYPE_CHECKING +from scepter.modules.utils.import_utils import LazyImportModule + +if TYPE_CHECKING: + from scepter.modules.annotator.base_annotator import GeneralAnnotator + from scepter.modules.annotator.canny import CannyAnnotator + from scepter.modules.annotator.color import ColorAnnotator + from scepter.modules.annotator.degradation import DegradationAnnotator + from scepter.modules.annotator.doodle import DoodleAnnotator + from scepter.modules.annotator.gray import GrayAnnotator + from scepter.modules.annotator.hed import HedAnnotator + from scepter.modules.annotator.identity import IdentityAnnotator + from scepter.modules.annotator.informative_drawing import ( + InfoDrawAnimeAnnotator, InfoDrawContourAnnotator, + InfoDrawOpenSketchAnnotator) + from scepter.modules.annotator.inpainting import InpaintingAnnotator + from scepter.modules.annotator.invert import InvertAnnotator + from scepter.modules.annotator.midas_op import MidasDetector + from scepter.modules.annotator.mlsd_op import MLSDdetector + from scepter.modules.annotator.openpose import OpenposeAnnotator + from scepter.modules.annotator.outpainting import OutpaintingAnnotator, OutpaintingResize + from scepter.modules.annotator.pidinet import PiDiAnnotator + from scepter.modules.annotator.segmentation import ESAMAnnotator + from scepter.modules.annotator.sketch import SketchAnnotator + from scepter.modules.annotator.lama import LamaAnnotator +else: + _import_structure = { + 'base_annotator': ['GeneralAnnotator'], + 'canny': ['CannyAnnotator'], + 'color': ['ColorAnnotator'], + 'degradation': ['DegradationAnnotator'], + 'doodle': ['DoodleAnnotator'], + 'gray': ['GrayAnnotator'], + 'hed': ['HedAnnotator'], + 'identity': ['IdentityAnnotator'], + 'informative_drawing': ['InfoDrawAnimeAnnotator', + 'InfoDrawContourAnnotator', + 'InfoDrawOpenSketchAnnotator'], + 'inpainting': ['InpaintingAnnotator'], + 'invert': ['InvertAnnotator'], + 'midas_op': ['MidasDetector'], + 'mlsd_op': ['MLSDdetector'], + 'openpose': ['OpenposeAnnotator'], + 'outpainting': ['OutpaintingAnnotator', 'OutpaintingResize'], + 'pidinet': ['PiDiAnnotator'], + 'segmentation': ['ESAMAnnotator'], + 'sketch': ['SketchAnnotator'], + 'lama': ['LamaAnnotator'], + } + + import sys + sys.modules[__name__] = LazyImportModule( + __name__, + globals()['__file__'], + _import_structure, + module_spec=__spec__, + extra_objects={}, + ) diff --git a/scepter/modules/annotator/hed.py b/scepter/modules/annotator/hed.py index 5956be7..bc68586 100644 --- a/scepter/modules/annotator/hed.py +++ b/scepter/modules/annotator/hed.py @@ -114,7 +114,7 @@ class HedAnnotator(BaseAnnotator, metaclass=ABCMeta): pretrained_model = cfg.get('PRETRAINED_MODEL', None) if pretrained_model: with FS.get_from(pretrained_model, wait_finish=True) as local_path: - self.netNetwork.load_state_dict(torch.load(local_path)) + self.netNetwork.load_state_dict(torch.load(local_path, weights_only=True)) @torch.no_grad() @torch.inference_mode() diff --git a/scepter/modules/annotator/informative_drawing.py b/scepter/modules/annotator/informative_drawing.py index 3303d09..532c888 100644 --- a/scepter/modules/annotator/informative_drawing.py +++ b/scepter/modules/annotator/informative_drawing.py @@ -120,7 +120,7 @@ class InfoDrawContourAnnotator(BaseAnnotator, metaclass=ABCMeta): self.model = ContourInference(input_nc, output_nc, n_residual_blocks, sigmoid) with FS.get_from(pretrained_model, wait_finish=True) as local_path: - self.model.load_state_dict(torch.load(local_path)) + self.model.load_state_dict(torch.load(local_path, weights_only=True)) self.model = self.model.eval().requires_grad_(False).to(we.device_id) @torch.no_grad() diff --git a/scepter/modules/annotator/midas/base_model.py b/scepter/modules/annotator/midas/base_model.py index c51ac55..2f99b8e 100644 --- a/scepter/modules/annotator/midas/base_model.py +++ b/scepter/modules/annotator/midas/base_model.py @@ -10,7 +10,7 @@ class BaseModel(torch.nn.Module): Args: path (str): file path """ - parameters = torch.load(path, map_location=torch.device('cpu')) + parameters = torch.load(path, map_location=torch.device('cpu'), weights_only=True) if 'optimizer' in parameters: parameters = parameters['model'] diff --git a/scepter/modules/annotator/mlsd_op.py b/scepter/modules/annotator/mlsd_op.py index 1c778c8..b7e1f12 100644 --- a/scepter/modules/annotator/mlsd_op.py +++ b/scepter/modules/annotator/mlsd_op.py @@ -29,7 +29,7 @@ class MLSDdetector(BaseAnnotator, metaclass=ABCMeta): pretrained_model = cfg.get('PRETRAINED_MODEL', None) if pretrained_model: with FS.get_from(pretrained_model, wait_finish=True) as local_path: - model.load_state_dict(torch.load(local_path), strict=True) + model.load_state_dict(torch.load(local_path, weights_only=True), strict=True) self.model = model.eval() self.thr_v = cfg.get('THR_V', 0.1) self.thr_d = cfg.get('THR_D', 0.1) diff --git a/scepter/modules/annotator/openpose.py b/scepter/modules/annotator/openpose.py index 8996afa..4f0135c 100644 --- a/scepter/modules/annotator/openpose.py +++ b/scepter/modules/annotator/openpose.py @@ -423,7 +423,7 @@ class Hand(object): self.model = handpose_model() if torch.cuda.is_available(): self.model = self.model.to(device) - model_dict = transfer(self.model, torch.load(model_path)) + model_dict = transfer(self.model, torch.load(model_path, weights_only=True)) self.model.load_state_dict(model_dict) self.model.eval() self.device = device @@ -503,7 +503,7 @@ class Body(object): self.model = bodypose_model() if torch.cuda.is_available(): self.model = self.model.to(device) - model_dict = transfer(self.model, torch.load(model_path)) + model_dict = transfer(self.model, torch.load(model_path, weights_only=True)) self.model.load_state_dict(model_dict) self.model.eval() self.device = device diff --git a/scepter/modules/annotator/pidinet.py b/scepter/modules/annotator/pidinet.py index bac3f80..a06ba4a 100644 --- a/scepter/modules/annotator/pidinet.py +++ b/scepter/modules/annotator/pidinet.py @@ -882,7 +882,7 @@ class PiDiAnnotator(BaseAnnotator, metaclass=ABCMeta): if pretrained_model: with FS.get_from(pretrained_model, wait_finish=True) as local_path: state = torch.load(local_path, - map_location='cpu')['state_dict'] + map_location='cpu', weights_only=True)['state_dict'] if vanilla_cnn: state = convert_pidinet(state, 'carv4') state = { diff --git a/scepter/modules/annotator/segmentation.py b/scepter/modules/annotator/segmentation.py index 073d9e3..ed884bf 100644 --- a/scepter/modules/annotator/segmentation.py +++ b/scepter/modules/annotator/segmentation.py @@ -10,7 +10,11 @@ import torchvision.transforms as T from PIL import Image from pycocotools import mask as mask_utils from scipy import ndimage -from sklearn.cluster import KMeans +try: + from sklearn.cluster import KMeans +except: + import warnings + warnings.warn("ignore sklearn import, please pip install scikit-learn.") from torchvision.ops.boxes import batched_nms from scepter.modules.annotator.base_annotator import BaseAnnotator diff --git a/scepter/modules/annotator/sketch.py b/scepter/modules/annotator/sketch.py index 1732ab6..53fb65f 100644 --- a/scepter/modules/annotator/sketch.py +++ b/scepter/modules/annotator/sketch.py @@ -86,7 +86,7 @@ class SketchAnnotator(BaseAnnotator, metaclass=ABCMeta): std=0.0858381272736797).eval() if pretrained_model: with FS.get_from(pretrained_model, wait_finish=True) as local_path: - state = torch.load(local_path, map_location='cpu') + state = torch.load(local_path, map_location='cpu', weights_only=True) self.model.load_state_dict(state) @torch.no_grad() diff --git a/scepter/modules/data/__init__.py b/scepter/modules/data/__init__.py index 198e24f..18ea630 100644 --- a/scepter/modules/data/__init__.py +++ b/scepter/modules/data/__init__.py @@ -1,4 +1,21 @@ # -*- coding: utf-8 -*- # Copyright (c) Alibaba, Inc. and its affiliates. +from typing import TYPE_CHECKING +from scepter.modules.utils.import_utils import LazyImportModule -from scepter.modules.data import dataset, sampler + +if TYPE_CHECKING: + from scepter.modules.data import dataset, sampler +else: + _import_structure = { + 'data': ['dataset', 'sampler'] + } + + import sys + sys.modules[__name__] = LazyImportModule( + __name__, + globals()['__file__'], + _import_structure, + module_spec=__spec__, + extra_objects={}, + ) diff --git a/scepter/modules/data/dataset/__init__.py b/scepter/modules/data/dataset/__init__.py index 347f0c2..a935920 100644 --- a/scepter/modules/data/dataset/__init__.py +++ b/scepter/modules/data/dataset/__init__.py @@ -1,12 +1,35 @@ # -*- coding: utf-8 -*- # Copyright (c) Alibaba, Inc. and its affiliates. +from typing import TYPE_CHECKING +from scepter.modules.utils.import_utils import LazyImportModule -from scepter.modules.data.dataset.base_dataset import BaseDataset -from scepter.modules.data.dataset.dataset import (Image2ImageDataset, - ImageClassifyPublicDataset, - ImageTextPairDataset, - Text2ImageDataset) -from scepter.modules.data.dataset.ms_dataset import ( - ImageTextPairFolderDataset, ImageTextPairMSDataset) -from scepter.modules.data.dataset.registry import DATASETS -from scepter.modules.data.dataset.video_gen_dataset import VideoGenDataset \ No newline at end of file + +if TYPE_CHECKING: + from scepter.modules.data.dataset.base_dataset import BaseDataset + from scepter.modules.data.dataset.dataset import (Image2ImageDataset, + ImageClassifyPublicDataset, + ImageTextPairDataset, + Text2ImageDataset) + from scepter.modules.data.dataset.ms_dataset import ( + ImageTextPairFolderDataset, ImageTextPairMSDataset) + from scepter.modules.data.dataset.registry import DATASETS + from scepter.modules.data.dataset.video_gen_dataset import VideoGenDataset +else: + _import_structure = { + 'base_dataset': ['BaseDataset'], + 'dataset': ['Image2ImageDataset', 'ImageClassifyPublicDataset', + 'ImageTextPairDataset', 'Text2ImageDataset'], + 'ms_dataset': ['ImageTextPairFolderDataset', + 'ImageTextPairMSDataset'], + 'registry': ['DATASETS'], + 'video_gen_dataset': ['VideoGenDataset'] + } + + import sys + sys.modules[__name__] = LazyImportModule( + __name__, + globals()['__file__'], + _import_structure, + module_spec=__spec__, + extra_objects={}, + ) diff --git a/scepter/modules/data/dataset/base_dataset.py b/scepter/modules/data/dataset/base_dataset.py index e35c805..0ddec6e 100644 --- a/scepter/modules/data/dataset/base_dataset.py +++ b/scepter/modules/data/dataset/base_dataset.py @@ -82,7 +82,7 @@ class BaseDataset(Dataset, metaclass=ABCMeta): overwrite=False) self.worker_id = worker_id self.logger = self.worker_logger - self.local_we["seed"] += (worker_id + we.rank) + self.local_we["seed"] += (worker_id + self.local_we['rank'] * 1234) self.seed = self.local_we["seed"] we.set_env(self.local_we) diff --git a/scepter/modules/data/dataset/dataset.py b/scepter/modules/data/dataset/dataset.py index 1d63e97..ff5c62c 100644 --- a/scepter/modules/data/dataset/dataset.py +++ b/scepter/modules/data/dataset/dataset.py @@ -4,6 +4,7 @@ import numbers import os import sys +import copy from collections.abc import Iterable import numpy as np diff --git a/scepter/modules/data/dataset/ms_dataset.py b/scepter/modules/data/dataset/ms_dataset.py index 45716ed..9ec750e 100644 --- a/scepter/modules/data/dataset/ms_dataset.py +++ b/scepter/modules/data/dataset/ms_dataset.py @@ -386,6 +386,11 @@ class ImageTextPairMSDatasetForACE(BaseDataset): 'description': 'The keywords sign you want to add, which is like <{HIGHLIGHT_KEYWORDS}{KEYWORDS_SIGN}>' }, + 'ALIGN_SIZE': { + 'value': False, + 'description': + 'Whether ensure the size align between the source image and target image.' + }, 'OUTPUT_SIZE': { 'value': None, @@ -414,6 +419,8 @@ class ImageTextPairMSDatasetForACE(BaseDataset): self.replace_keywords = cfg.get('HIGHLIGHT_KEYWORDS', '') self.keywords_sign = cfg.get('KEYWORDS_SIGN', '') self.add_indicator = cfg.get('ADD_INDICATOR', False) + + self.align_size = cfg.get('ALIGN_SIZE', False) # Use modelscope dataset if not ms_dataset_name: raise ValueError( @@ -492,7 +499,7 @@ class ImageTextPairMSDatasetForACE(BaseDataset): tar_image_path, cvt_type='RGB') src_image = self.image_preprocess(src_image) - tar_image = self.image_preprocess(tar_image) + tar_image = self.image_preprocess(tar_image, size = src_image.shape[:2] if self.align_size else None) tar_image = self.transforms(tar_image) src_image = self.transforms(src_image) @@ -501,13 +508,13 @@ class ImageTextPairMSDatasetForACE(BaseDataset): if self.add_indicator: if '{image}' not in prompt: prompt = '{image}, ' + prompt - return { - 'edit_image': [src_image], - 'edit_image_mask': [src_mask], + 'src_image_list': [src_image], + 'src_mask_list': [src_mask], 'image': tar_image, 'image_mask': tar_mask, 'prompt': [prompt], + 'edit_id': [0] } def load_image(self, prefix, img_path, cvt_type=None): diff --git a/scepter/modules/data/dataset/registry.py b/scepter/modules/data/dataset/registry.py index 3906d41..33ddaaf 100644 --- a/scepter/modules/data/dataset/registry.py +++ b/scepter/modules/data/dataset/registry.py @@ -337,8 +337,12 @@ def build_dataset_config(cfg, registry, logger=None, *args, **kwargs): f'registry must be type Registry, got {type(registry)}') cfg = deep_copy(cfg) - req_type = cfg.get('NAME') + + from scepter.modules.utils.import_utils import LazyImportModule + sig = (registry.name.upper(), req_type) + LazyImportModule.import_module(sig) + if isinstance(req_type, str): req_type_entry = registry.get(req_type) if req_type_entry is None: diff --git a/scepter/modules/data/dataset/video_gen_dataset.py b/scepter/modules/data/dataset/video_gen_dataset.py index 16a5dd1..791b890 100644 --- a/scepter/modules/data/dataset/video_gen_dataset.py +++ b/scepter/modules/data/dataset/video_gen_dataset.py @@ -1,44 +1,46 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. import io +import os import random import sys -import os import warnings -import torch import numpy as np +import torch from tqdm import tqdm -from scepter.modules.utils.distribute import we from scepter.modules.data.dataset import DATASETS, BaseDataset +from scepter.modules.utils.distribute import we from scepter.modules.utils.file_system import FS try: import decord - decord.bridge.set_bridge("torch") + decord.bridge.set_bridge('torch') except ImportError: warnings.warn( - "The `decord` package is required for loading the video dataset. Install with `pip install decord`" + 'The `decord` package is required for loading the video dataset. Install with `pip install decord`' ) @DATASETS.register_class() class VideoGenDataset(BaseDataset): - def __init__(self, cfg, logger = None): + def __init__(self, cfg, logger=None): super().__init__(cfg, logger=logger) self.prompt_prefix = cfg.get('PROMPT_PREFIX', '') self.path_prefix = cfg.get('PATH_PREFIX', '') self.p_zero = cfg.get('P_ZERO', 0.0) - self.max_num_frames = cfg.get("NUM_FRAMES", 49) - self.fps = cfg.get("FPS", 8) - self.height = cfg.get("HEIGHT", 480) - self.width = cfg.get("WIDTH", 720) - self.skip_frames_start = cfg.get("SKIP_FRAMES_START", 0) - self.skip_frames_end = cfg.get("SKIP_FRAMES_END", 0) + self.max_num_frames = cfg.get('NUM_FRAMES', 49) + self.fps = cfg.get('FPS', 8) + self.height = cfg.get('HEIGHT', 480) + self.width = cfg.get('WIDTH', 720) + self.skip_frames_start = cfg.get('SKIP_FRAMES_START', 0) + self.skip_frames_end = cfg.get('SKIP_FRAMES_END', 0) self.data_type = cfg.get('DATA_TYPE', 't2v') def worker_init_fn(self, worker_id, num_workers=1): super().worker_init_fn(worker_id, num_workers=num_workers) - randseed = np.random.randint(0, 2 ** 32 - num_workers - 1) + randseed = np.random.randint(0, 2**32 - num_workers - 1) workerseed = randseed + worker_id random.seed(workerseed) np.random.seed(workerseed) @@ -46,7 +48,9 @@ class VideoGenDataset(BaseDataset): def _preprocess_video_data(self, video_path): with FS.get_object(video_path) as video_data: - video_reader = decord.VideoReader(io.BytesIO(video_data), width=self.width, height=self.height) + video_reader = decord.VideoReader(io.BytesIO(video_data), + width=self.width, + height=self.height) video_num_frames = len(video_reader) start_frame = min(self.skip_frames_start, video_num_frames) @@ -54,13 +58,16 @@ class VideoGenDataset(BaseDataset): if end_frame <= start_frame: frames = video_reader.get_batch([start_frame]) elif end_frame - start_frame <= self.max_num_frames: - frames = video_reader.get_batch(list(range(start_frame, end_frame))) + frames = video_reader.get_batch(list(range(start_frame, + end_frame))) else: - indices = list(range(start_frame, end_frame, (end_frame - start_frame) // self.max_num_frames)) + indices = list( + range(start_frame, end_frame, + (end_frame - start_frame) // self.max_num_frames)) frames = video_reader.get_batch(indices) # Ensure that we don't go over the limit - frames = frames[: self.max_num_frames] + frames = frames[:self.max_num_frames] selected_num_frames = frames.shape[0] # Choose first (4k + 1) frames as this is how many is required by the VAE @@ -73,14 +80,16 @@ class VideoGenDataset(BaseDataset): # Training transforms frames = frames.float().div_(127.5).sub_(1.) - frames = frames.permute(3, 0, 1, 2).contiguous() # [C, F, H, W] + frames = frames.permute(3, 0, 1, 2).contiguous() # [C, F, H, W] return frames def _parse_index(self, index): meta = dict() for key, value in zip(index[-1], index[:-1]): - if key in ['oss_key', 'path', 'video_path']: + if key in ['oss_key', 'path', 'video_path', 'target_video_path']: meta['video_path'] = value + elif key in ['source_video_path', 'src_video_path']: + meta['src_video_path'] = value elif key in ['prompt', 'caption', 'text']: meta['prompt'] = value elif key in ['width', 'height']: @@ -104,8 +113,13 @@ class VideoGenDataset(BaseDataset): 'prompt': prompt, 'meta': meta, } - if self.data_type == 'i2v': + if 'i2v' in self.data_type: item['image'] = item['video'][:, :1, :, :] + if 'v2v' in self.data_type: + src_video_path = os.path.join(self.path_prefix, + meta.get('src_video_path', '')) + src_video = self._preprocess_video_data(src_video_path) + item['src_video'] = src_video return item def __len__(self): @@ -122,7 +136,6 @@ class VideoGenDataset(BaseDataset): return collect - @DATASETS.register_class() class VideoGenDatasetOTF(VideoGenDataset): def __init__(self, cfg, logger=None): @@ -135,8 +148,11 @@ class VideoGenDatasetOTF(VideoGenDataset): from scepter.modules.model.registry import MODELS model_cfg = cfg.get('MODEL', None) if model_cfg is not None: - self.model = MODELS.build(cfg.MODEL, logger=logger).eval().requires_grad_(False).to(we.device_id) - self.items = self.parse_data(self.data_file, self.delimiter, self.fields) + self.model = MODELS.build( + cfg.MODEL, + logger=logger).eval().requires_grad_(False).to(we.device_id) + self.items = self.parse_data(self.data_file, self.delimiter, + self.fields) if self.use_num and self.use_num > 0: self.items = self.items[:self.use_num] self.data = self.encode(self.items) @@ -169,11 +185,13 @@ class VideoGenDatasetOTF(VideoGenDataset): return items def encode(self, items): - self.logger.info("Start to encode video data [{}]!".format(len(items))) + self.logger.info('Start to encode video data [{}]!'.format(len(items))) for item in tqdm(items): - video_path = os.path.join(self.path_prefix, item.get('video_path', '')) + video_path = os.path.join(self.path_prefix, + item.get('video_path', '')) video = self._preprocess_video_data(video_path) - latent = self.model.encode_first_stage(video.unsqueeze(0).to(we.device_id)).squeeze(0) + latent = self.model.encode_first_stage( + video.unsqueeze(0).to(we.device_id)).squeeze(0) item['video_latent'] = latent.detach().cpu() item['video'] = video if self.data_type == 'i2v': @@ -181,4 +199,4 @@ class VideoGenDatasetOTF(VideoGenDataset): return items def _get(self, index): - return self.data[index % self.real_number] \ No newline at end of file + return self.data[index % self.real_number] diff --git a/scepter/modules/data/sampler/__init__.py b/scepter/modules/data/sampler/__init__.py index 428ec1b..50d3380 100644 --- a/scepter/modules/data/sampler/__init__.py +++ b/scepter/modules/data/sampler/__init__.py @@ -1,9 +1,31 @@ # -*- coding: utf-8 -*- # Copyright (c) Alibaba, Inc. and its affiliates. +from typing import TYPE_CHECKING +from scepter.modules.utils.import_utils import LazyImportModule -from scepter.modules.data.sampler.base_sampler import BaseSampler -from scepter.modules.data.sampler.registry import SAMPLERS -from scepter.modules.data.sampler.sampler import ( - EvalDistributedSampler, LoopSampler, MixtureOfSamplers, - MultiFoldDistributedSampler, MultiLevelBatchSampler, - MultiLevelBatchSamplerMultiSource, ResolutionBatchSampler) + +if TYPE_CHECKING: + from scepter.modules.data.sampler.base_sampler import BaseSampler + from scepter.modules.data.sampler.registry import SAMPLERS + from scepter.modules.data.sampler.sampler import ( + EvalDistributedSampler, LoopSampler, MixtureOfSamplers, + MultiFoldDistributedSampler, MultiLevelBatchSampler, + MultiLevelBatchSamplerMultiSource, ResolutionBatchSampler) +else: + _import_structure = { + 'base_sampler': ['BaseSampler'], + 'registry': ['SAMPLERS'], + 'sampler': ['EvalDistributedSampler', 'LoopSampler', + 'MixtureOfSamplers', 'MultiFoldDistributedSampler', + 'MultiLevelBatchSampler', 'MultiLevelBatchSamplerMultiSource', + 'ResolutionBatchSampler'] + } + + import sys + sys.modules[__name__] = LazyImportModule( + __name__, + globals()['__file__'], + _import_structure, + module_spec=__spec__, + extra_objects={}, + ) diff --git a/scepter/modules/data/sampler/registry.py b/scepter/modules/data/sampler/registry.py index 4a3324b..4c67a76 100644 --- a/scepter/modules/data/sampler/registry.py +++ b/scepter/modules/data/sampler/registry.py @@ -35,8 +35,12 @@ def build_sampler_config(cfg, registry, logger=None, **kwargs): f'registry must be type Registry, got {type(registry)}') cfg = deep_copy(cfg) - req_type = cfg.get('NAME') + + from scepter.modules.utils.import_utils import LazyImportModule + sig = (registry.name.upper(), req_type) + LazyImportModule.import_module(sig) + if isinstance(req_type, str): req_type_entry = registry.get(req_type) if req_type_entry is None: diff --git a/scepter/modules/data/utils/__init__.py b/scepter/modules/data/utils/__init__.py index 75399c0..c67fbac 100644 --- a/scepter/modules/data/utils/__init__.py +++ b/scepter/modules/data/utils/__init__.py @@ -1,4 +1,21 @@ # -*- coding: utf-8 -*- # Copyright (c) Alibaba, Inc. and its affiliates. +from typing import TYPE_CHECKING +from scepter.modules.utils.import_utils import LazyImportModule -from scepter.modules.data.utils.data_bucket import BucketManager + +if TYPE_CHECKING: + from scepter.modules.data.utils.data_bucket import BucketManager +else: + _import_structure = { + 'data_bucket': ['BucketManager'] + } + + import sys + sys.modules[__name__] = LazyImportModule( + __name__, + globals()['__file__'], + _import_structure, + module_spec=__spec__, + extra_objects={}, + ) diff --git a/scepter/modules/inference/__init__.py b/scepter/modules/inference/__init__.py index d442d9c..dad520f 100644 --- a/scepter/modules/inference/__init__.py +++ b/scepter/modules/inference/__init__.py @@ -1,3 +1,39 @@ # -*- coding: utf-8 -*- # Copyright (c) Alibaba, Inc. and its affiliates. -from scepter.modules.inference.diffusion_inference import DiffusionInference +from typing import TYPE_CHECKING +from scepter.modules.utils.import_utils import LazyImportModule + + +if TYPE_CHECKING: + from scepter.modules.inference.diffusion_inference import DiffusionInference + from scepter.modules.inference.ace_inference import ACEInference + from scepter.modules.inference.cogvideox_inference import CogVideoXInference + from scepter.modules.inference.control_inference import ControlInference + from scepter.modules.inference.flux_inference import FluxInference + 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.inference.tuner_inference import TunerInference +else: + _import_structure = { + 'diffusion_inference': ['DiffusionInference'], + 'ace_inference': ['ACEInference'], + 'cogvideox_inference': ['CogVideoXInference'], + 'control_inference': ['ControlInference'], + 'flux_inference': ['FluxInference'], + 'largen_inference': ['LargenInference'], + 'pixart_inference': ['PixArtInference'], + 'sd3_inference': ['SD3Inference'], + 'stylebooth_inference': ['StyleboothInference'], + 'tuner_inference': ['TunerInference'] + } + + import sys + sys.modules[__name__] = LazyImportModule( + __name__, + globals()['__file__'], + _import_structure, + module_spec=__spec__, + extra_objects={}, + ) diff --git a/scepter/modules/inference/cogvideox_inference.py b/scepter/modules/inference/cogvideox_inference.py index a740361..b43ddb9 100644 --- a/scepter/modules/inference/cogvideox_inference.py +++ b/scepter/modules/inference/cogvideox_inference.py @@ -11,7 +11,6 @@ import torch from scepter.modules.utils.file_system import FS from scepter.modules.utils.distribute import we from scepter.modules.model.backbone.cogvideox.utils import get_3d_rotary_pos_embed, get_resize_crop_region_for_grid - from .diffusion_inference import DiffusionInference, get_model from .tuner_inference import TunerInference @@ -27,9 +26,13 @@ class CogVideoXInference(DiffusionInference): @torch.no_grad() def decode_first_stage(self, latents): - latents = latents.permute(0, 2, 1, 3, 4) - latents = 1 / self.first_stage_model['paras']['scaling_factor_image'] * latents - frames = get_model(self.first_stage_model).decode(latents) + _, dtype = self.get_function_info(self.first_stage_model, 'decode') + with torch.autocast('cuda', + enabled=dtype in ('bfloat16'), + dtype=getattr(torch, dtype)): + latents = latents.permute(0, 2, 1, 3, 4) + latents = 1 / self.first_stage_model['paras']['scaling_factor_image'] * latents + frames = get_model(self.first_stage_model).decode(latents) return frames def _prepare_rotary_positional_embeddings( @@ -121,8 +124,9 @@ class CogVideoXInference(DiffusionInference): ) function_name, dtype = self.get_function_info( self.diffusion_model) + with torch.autocast('cuda', - enabled=dtype=='bfloat16', + enabled=dtype in ('float16', 'bfloat16'), dtype=getattr(torch, dtype)): solver_sample = value_input.get('sample', 'ddim') sample_steps = value_input.get('sample_steps', 50) @@ -143,7 +147,6 @@ class CogVideoXInference(DiffusionInference): }], steps=sample_steps, show_progress=True, - use_dynamic_cfg=True, guide_scale=guide_scale, guide_rescale=guide_rescale, return_intermediate=None, @@ -151,7 +154,6 @@ class CogVideoXInference(DiffusionInference): self.dynamic_unload(self.diffusion_model, 'diffusion_model', skip_loaded=True) - self.dynamic_load(self.first_stage_model, 'first_stage_model') x_samples = self.decode_first_stage(latent).float() # [B, C, F, H, W] self.dynamic_unload(self.first_stage_model, diff --git a/scepter/modules/inference/diffusion_inference.py b/scepter/modules/inference/diffusion_inference.py index cc8c0b8..1715245 100644 --- a/scepter/modules/inference/diffusion_inference.py +++ b/scepter/modules/inference/diffusion_inference.py @@ -97,7 +97,7 @@ class DiffusionInference(): if 'weights_only' in torch.load.__code__.co_varnames: sd = torch.load(local_path, map_location='cpu', weights_only=True) else: - sd = torch.load(local_path, map_location='cpu') + sd = torch.load(local_path, map_location='cpu', weights_only=True) first_stage_model_path = os.path.join( os.path.dirname(local_path), 'first_stage_model.pth') cond_stage_model_path = os.path.join( @@ -203,7 +203,7 @@ class DiffusionInference(): from safetensors.torch import load_file as load_safetensors sd = load_safetensors(path) else: - sd = torch.load(path, map_location='cpu') + sd = torch.load(path, map_location='cpu', weights_only=True) new_sd = OrderedDict() for k, v in sd.items(): @@ -230,16 +230,22 @@ class DiffusionInference(): def load(self, module): if module['device'] == 'offline': - if module['cfg'].NAME in MODELS.class_map: + from scepter.modules.utils.import_utils import LazyImportModule + if (LazyImportModule.get_module_type(('MODELS', module['cfg'].NAME)) or + module['cfg'].NAME in MODELS.class_map): model = MODELS.build(module['cfg'], logger=self.logger).eval() - elif module['cfg'].NAME in BACKBONES.class_map: + elif (LazyImportModule.get_module_type(('BACKBONES', module['cfg'].NAME)) or + module['cfg'].NAME in BACKBONES.class_map): model = BACKBONES.build(module['cfg'], logger=self.logger).eval() - elif module['cfg'].NAME in EMBEDDERS.class_map: + elif (LazyImportModule.get_module_type(('EMBEDDERS', module['cfg'].NAME)) or + module['cfg'].NAME in EMBEDDERS.class_map): model = EMBEDDERS.build(module['cfg'], logger=self.logger).eval() else: raise NotImplementedError + if 'DTYPE' in module['cfg'] and module['cfg']['DTYPE'] is not None: + model = model.to(getattr(torch, module['cfg'].DTYPE)) if module['cfg'].get('RELOAD_MODEL', None): self.init_from_ckpt(module['cfg'].RELOAD_MODEL, model) module['model'] = model @@ -268,8 +274,9 @@ class DiffusionInference(): module['device'] = 'cpu' else: module['device'] = 'offline' - torch.cuda.empty_cache() - torch.cuda.ipc_collect() + if torch.cuda.is_available(): + torch.cuda.empty_cache() + torch.cuda.ipc_collect() return module def dynamic_load(self, module=None, name=''): diff --git a/scepter/modules/inference/largen_inference.py b/scepter/modules/inference/largen_inference.py index 86c4848..ddb52ba 100644 --- a/scepter/modules/inference/largen_inference.py +++ b/scepter/modules/inference/largen_inference.py @@ -43,7 +43,7 @@ class LargenInference(DiffusionInference): if 'weights_only' in torch.load.__code__.co_varnames: sd = torch.load(local_path, map_location='cpu', weights_only=True) else: - sd = torch.load(local_path, map_location='cpu') + sd = torch.load(local_path, map_location='cpu', weights_only=True) if 'model' in sd: sd = sd['model'] diff --git a/scepter/modules/inference/tuner_inference.py b/scepter/modules/inference/tuner_inference.py index db73795..6d30a41 100644 --- a/scepter/modules/inference/tuner_inference.py +++ b/scepter/modules/inference/tuner_inference.py @@ -144,9 +144,9 @@ class TunerInference(): is_bin_file = True if os.path.isfile(bin_file): if 'weights_only' in torch.load.__code__.co_varnames: - state_dict = torch.load(bin_file, weights_only=True) + state_dict = torch.load(bin_file, weights_only=True, map_location="cpu") else: - state_dict = torch.load(bin_file) + state_dict = torch.load(bin_file, map_location="cpu") elif os.path.isfile(safe_file): is_bin_file = False from safetensors.torch import \ diff --git a/scepter/modules/model/__init__.py b/scepter/modules/model/__init__.py index bf33490..4494010 100644 --- a/scepter/modules/model/__init__.py +++ b/scepter/modules/model/__init__.py @@ -1,5 +1,23 @@ # -*- coding: utf-8 -*- # Copyright (c) Alibaba, Inc. and its affiliates. +from typing import TYPE_CHECKING +from scepter.modules.utils.import_utils import LazyImportModule -from scepter.modules.model import (backbone, embedder, head, loss, metric, - neck, network, tokenizer, tuner, diffusion) + +if TYPE_CHECKING: + from scepter.modules.model import (backbone, embedder, head, loss, metric, + neck, network, tokenizer, tuner, diffusion) +else: + _import_structure = { + 'model': ['backbone', 'embedder', 'head', 'loss', 'metric', + 'neck', 'network', 'tokenizer', 'tuner', 'diffusion'] + } + + import sys + sys.modules[__name__] = LazyImportModule( + __name__, + globals()['__file__'], + _import_structure, + module_spec=__spec__, + extra_objects={}, + ) diff --git a/scepter/modules/model/backbone/__init__.py b/scepter/modules/model/backbone/__init__.py index 44b6480..deafbb9 100644 --- a/scepter/modules/model/backbone/__init__.py +++ b/scepter/modules/model/backbone/__init__.py @@ -1,4 +1,23 @@ # -*- coding: utf-8 -*- # Copyright (c) Alibaba, Inc. and its affiliates. -from scepter.modules.model.backbone import (ace, autoencoder, flux, image, cogvideox, - mmdit, pixart, unet, utils, video) +from typing import TYPE_CHECKING +from scepter.modules.utils.import_utils import LazyImportModule + + +if TYPE_CHECKING: + from scepter.modules.model.backbone import (ace, autoencoder, flux, image, cogvideox, + mmdit, pixart, unet, utils, video) +else: + _import_structure = { + 'backbone': ['ace', 'autoencoder', 'flux', 'image', 'cogvideox', + 'mmdit', 'pixart', 'unet', 'utils', 'video'] + } + + import sys + sys.modules[__name__] = LazyImportModule( + __name__, + globals()['__file__'], + _import_structure, + module_spec=__spec__, + extra_objects={}, + ) diff --git a/scepter/modules/model/backbone/ace/ace.py b/scepter/modules/model/backbone/ace/ace.py index 0873413..3e8253e 100644 --- a/scepter/modules/model/backbone/ace/ace.py +++ b/scepter/modules/model/backbone/ace/ace.py @@ -151,7 +151,7 @@ class ACE(BaseModel): 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') + model = torch.load(local_path, map_location='cpu', weights_only=True) if 'state_dict' in model: model = model['state_dict'] new_ckpt = OrderedDict() diff --git a/scepter/modules/model/backbone/ace_plus/__init__.py b/scepter/modules/model/backbone/ace_plus/__init__.py new file mode 100644 index 0000000..cc26a06 --- /dev/null +++ b/scepter/modules/model/backbone/ace_plus/__init__.py @@ -0,0 +1,2 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. diff --git a/scepter/modules/model/backbone/ace_plus/ace_plus.py b/scepter/modules/model/backbone/ace_plus/ace_plus.py new file mode 100644 index 0000000..bcec907 --- /dev/null +++ b/scepter/modules/model/backbone/ace_plus/ace_plus.py @@ -0,0 +1,97 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +import torch +from torch.nn.utils.rnn import pad_sequence + +from einops import rearrange +from scepter.modules.model.backbone.flux import FluxMR +from scepter.modules.model.registry import BACKBONES +from scepter.modules.utils.config import dict_to_yaml + + +@BACKBONES.register_class() +class FluxMRACEPlus(FluxMR): + def __init__(self, cfg, logger=None): + super().__init__(cfg, logger) + + def prepare_input(self, x, cond): + context, y = cond['context'], cond['y'] + batch_frames, batch_frames_ids = [], [] + for ix, shape, imask, ie, ie_mask in zip(x, cond['x_shapes'], + cond['x_mask'], cond['edit'], + cond['edit_mask']): + # unpack image from sequence + ix = ix[:, :shape[0] * shape[1]].view(-1, shape[0], shape[1]) + imask = torch.ones_like( + ix[[0], :, :]) if imask is None else imask.squeeze(0) + if len(ie) > 0: + ie = [iie.squeeze(0) for iie in ie] + ie_mask = [ + torch.ones( + (ix.shape[0] * 4, ix.shape[1], + ix.shape[2])) if iime is None else iime.squeeze(0) + for iime in ie_mask + ] + ie = torch.cat(ie, dim=-1) + ie_mask = torch.cat(ie_mask, dim=-1) + else: + ie, ie_mask = torch.zeros_like(ix).to(x), torch.ones_like( + imask).to(x) + ix = torch.cat([ix, ie, ie_mask], dim=0) + c, h, w = ix.shape + ix = rearrange(ix, + 'c (h ph) (w pw) -> (h w) (c ph pw)', + ph=2, + pw=2) + ix_id = torch.zeros(h // 2, w // 2, 3) + ix_id[..., 1] = ix_id[..., 1] + torch.arange(h // 2)[:, None] + ix_id[..., 2] = ix_id[..., 2] + torch.arange(w // 2)[None, :] + ix_id = rearrange(ix_id, 'h w c -> (h w) c') + batch_frames.append([ix]) + batch_frames_ids.append([ix_id]) + x_list, x_id_list, mask_x_list, x_seq_length = [], [], [], [] + for frames, frame_ids in zip(batch_frames, batch_frames_ids): + proj_frames = [] + for idx, one_frame in enumerate(frames): + one_frame = self.img_in(one_frame) + proj_frames.append(one_frame) + ix = torch.cat(proj_frames, dim=0) + if_id = torch.cat(frame_ids, dim=0) + x_list.append(ix) + x_id_list.append(if_id) + mask_x_list.append( + torch.ones(ix.shape[0]).to(ix.device, + non_blocking=True).bool()) + x_seq_length.append(ix.shape[0]) + # if len(x_list) < 1: import pdb;pdb.set_trace() + x = pad_sequence(tuple(x_list), batch_first=True) + x_ids = pad_sequence(tuple(x_id_list), batch_first=True).to( + x) # [b,pad_seq,2] pad (0.,0.) at dim2 + mask_x = pad_sequence(tuple(mask_x_list), batch_first=True) + # import pdb;pdb.set_trace() + if isinstance(context, list): + txt_list, mask_txt_list, y_list = [], [], [] + for sample_id, (ctx, yy) in enumerate(zip(context, y)): + txt_list.append(self.txt_in(ctx.to(x))) + mask_txt_list.append( + torch.ones(txt_list[-1].shape[0]).to( + ctx.device, non_blocking=True).bool()) + y_list.append(yy.to(x)) + txt = pad_sequence(tuple(txt_list), batch_first=True) + txt_ids = torch.zeros(txt.shape[0], txt.shape[1], 3).to(x) + mask_txt = pad_sequence(tuple(mask_txt_list), batch_first=True) + y = torch.cat(y_list, dim=0) + assert y.ndim == 2 and txt.ndim == 3 + else: + txt = self.txt_in(context) + txt_ids = torch.zeros(context.shape[0], context.shape[1], 3).to(x) + mask_txt = torch.ones(context.shape[0], context.shape[1]).to( + x.device, non_blocking=True).bool() + return x, x_ids, txt, txt_ids, y, mask_x, mask_txt, x_seq_length + + @staticmethod + def get_config_template(): + return dict_to_yaml('MODEL', + __class__.__name__, + FluxMRACEPlus.para_dict, + set_name=True) diff --git a/scepter/modules/model/backbone/cogvideox/cogvideox.py b/scepter/modules/model/backbone/cogvideox/cogvideox.py index dcb0fdf..67e70dd 100644 --- a/scepter/modules/model/backbone/cogvideox/cogvideox.py +++ b/scepter/modules/model/backbone/cogvideox/cogvideox.py @@ -48,6 +48,8 @@ class CogVideoXTransformer3DModel(BaseModel): Whether to flip the sin to cos in the time embedding. time_embed_dim (`int`, defaults to `512`): Output dimension of timestep embeddings. + ofs_embed_dim (`int`, defaults to `512`): + Output dimension of "ofs" embeddings used in CogVideoX-5b-I2B in version 1.5 text_embed_dim (`int`, defaults to `4096`): Input dimension of text embeddings from the text encoder. num_layers (`int`, defaults to `30`): @@ -98,6 +100,7 @@ class CogVideoXTransformer3DModel(BaseModel): flip_sin_to_cos = cfg.get("FLIP_SIN_TO_COS", True) freq_shift = cfg.get("FREQ_SHIFT", 0) time_embed_dim = cfg.get("TIME_EMBED_DIM", 512) + ofs_embed_dim = cfg.get("OFS_EMBED_DIM", None) # 1.5 text_embed_dim = cfg.get("TEXT_EMBED_DIM", 4096) num_layers = cfg.get("NUM_LAYERS", 30) dropout = cfg.get("DROPOUT", 0.0) @@ -106,6 +109,8 @@ class CogVideoXTransformer3DModel(BaseModel): sample_height = cfg.get("SAMPLE_HEIGHT", 60) sample_frames = cfg.get("SAMPLE_FRAMES", 49) patch_size = cfg.get("PATCH_SIZE", 2) + patch_size_t = cfg.get("PATCH_SIZE_T", None) + patch_bias = cfg.get("PATCH_BIAS", True) temporal_compression_ratio = cfg.get("TEMPORAL_COMPRESSION_RATIO", 4) max_text_seq_length = cfg.get("MAX_TEXT_SEQ_LENGTH", 226) activation_fn = cfg.get("ACTIVATION_FN", "gelu-approximate") @@ -119,6 +124,7 @@ class CogVideoXTransformer3DModel(BaseModel): self.gradient_checkpointing = cfg.get("GRADIENT_CHECKPOINTING", False) inner_dim = num_attention_heads * attention_head_dim self.patch_size = patch_size + self.patch_size_t = patch_size_t self.use_rotary_positional_embeddings = use_rotary_positional_embeddings if not use_rotary_positional_embeddings and use_learned_positional_embeddings: @@ -131,10 +137,11 @@ class CogVideoXTransformer3DModel(BaseModel): # 1. Patch embedding self.patch_embed = CogVideoXPatchEmbed( patch_size=patch_size, + patch_size_t=patch_size_t, in_channels=in_channels, embed_dim=inner_dim, text_embed_dim=text_embed_dim, - bias=True, + bias=patch_bias, sample_width=sample_width, sample_height=sample_height, sample_frames=sample_frames, @@ -147,10 +154,18 @@ class CogVideoXTransformer3DModel(BaseModel): ) self.embedding_dropout = nn.Dropout(dropout) - # 2. Time embeddings + # 2. Time embeddings and ofs embedding(Only CogVideoX1.5-5B I2V have) self.time_proj = Timesteps(inner_dim, flip_sin_to_cos, freq_shift) self.time_embedding = TimestepEmbedding(inner_dim, time_embed_dim, timestep_activation_fn) + self.ofs_proj = None + self.ofs_embedding = None + if ofs_embed_dim: + self.ofs_proj = Timesteps(ofs_embed_dim, flip_sin_to_cos, freq_shift) + self.ofs_embedding = TimestepEmbedding( + ofs_embed_dim, ofs_embed_dim, timestep_activation_fn + ) # same as time embeddings, for ofs + # 3. Define spatio-temporal transformers blocks self.transformer_blocks = nn.ModuleList( [ @@ -178,7 +193,15 @@ class CogVideoXTransformer3DModel(BaseModel): norm_eps=norm_eps, chunk_dim=1, ) - self.proj_out = nn.Linear(inner_dim, patch_size * patch_size * out_channels) + + if patch_size_t is None: + # For CogVideox 1.0 + output_dim = patch_size * patch_size * out_channels + else: + # For CogVideoX 1.5 + output_dim = patch_size * patch_size * patch_size_t * out_channels + + self.proj_out = nn.Linear(inner_dim, output_dim) def forward( self, @@ -186,6 +209,7 @@ class CogVideoXTransformer3DModel(BaseModel): t: Union[int, float, torch.LongTensor] = None, cond: torch.Tensor = None, timestep_cond: Optional[torch.Tensor] = None, + ofs: Optional[Union[int, float, torch.LongTensor]] = None, image_rotary_emb: Optional[Tuple[torch.Tensor, torch.Tensor]] = None, **kwargs ): @@ -208,6 +232,12 @@ class CogVideoXTransformer3DModel(BaseModel): t_emb = t_emb.to(dtype=encoder_hidden_states.dtype) emb = self.time_embedding(t_emb, timestep_cond) + if self.ofs_embedding is not None: + ofs_emb = self.ofs_proj(ofs) + ofs_emb = ofs_emb.to(dtype=hidden_states.dtype) + ofs_emb = self.ofs_embedding(ofs_emb) + emb = emb + ofs_emb + # 2. Patch embedding hidden_states = self.patch_embed(encoder_hidden_states, hidden_states) hidden_states = self.embedding_dropout(hidden_states) @@ -261,8 +291,16 @@ class CogVideoXTransformer3DModel(BaseModel): # - It is okay to `channels` use for CogVideoX-2b and CogVideoX-5b (number of input channels is equal to output channels) # - However, for CogVideoX-5b-I2V also takes concatenated input image latents (number of input channels is twice the output channels) p = self.patch_size - output = hidden_states.reshape(batch_size, num_frames, height // p, width // p, -1, p, p) - output = output.permute(0, 1, 4, 2, 5, 3, 6).flatten(5, 6).flatten(3, 4) + p_t = self.patch_size_t + + if p_t is None: + output = hidden_states.reshape(batch_size, num_frames, height // p, width // p, -1, p, p) + output = output.permute(0, 1, 4, 2, 5, 3, 6).flatten(5, 6).flatten(3, 4) + else: + output = hidden_states.reshape( + batch_size, (num_frames + p_t - 1) // p_t, height // p, width // p, -1, p_t, p, p + ) + output = output.permute(0, 1, 5, 4, 2, 6, 3, 7).flatten(6, 7).flatten(4, 5).flatten(1, 2) return output @@ -277,7 +315,7 @@ class CogVideoXTransformer3DModel(BaseModel): from safetensors.torch import load_file as load_safetensors ckpt = load_safetensors(local_model) else: - ckpt = torch.load(local_model, map_location='cpu') + ckpt = torch.load(local_model, map_location='cpu', weights_only=True) ckpt_all.update(ckpt) missing, unexpected = self.load_state_dict(ckpt_all, strict=False) if we.rank == 0: @@ -309,9 +347,9 @@ if __name__ == "__main__": FS.init_fs_client(file_sys) model = BACKBONES.build(cfg.DIFFUSION_MODEL, logger=get_logger()).eval().requires_grad_(False).to('cuda').to(torch.bfloat16) - hidden_states = torch.load(FS.get_from(cfg.HIDDEN_STATES)) - encoder_hidden_states = torch.load(FS.get_from(cfg.ENCODER_HIDDEN_STATES)) - timestep = torch.load(FS.get_from(cfg.TIMESTEP)) + hidden_states = torch.load(FS.get_from(cfg.HIDDEN_STATES), weights_only=True) + encoder_hidden_states = torch.load(FS.get_from(cfg.ENCODER_HIDDEN_STATES), weights_only=True) + timestep = torch.load(FS.get_from(cfg.TIMESTEP), weights_only=True) timestep_cond = None image_rotary_emb = None attention_kwargs = None diff --git a/scepter/modules/model/backbone/cogvideox/layers.py b/scepter/modules/model/backbone/cogvideox/layers.py index a1a1990..3a5c24f 100644 --- a/scepter/modules/model/backbone/cogvideox/layers.py +++ b/scepter/modules/model/backbone/cogvideox/layers.py @@ -178,6 +178,7 @@ class CogVideoXPatchEmbed(nn.Module): def __init__( self, patch_size: int = 2, + patch_size_t: Optional[int] = None, in_channels: int = 16, embed_dim: int = 1920, text_embed_dim: int = 4096, @@ -195,6 +196,7 @@ class CogVideoXPatchEmbed(nn.Module): super().__init__() self.patch_size = patch_size + self.patch_size_t = patch_size_t self.embed_dim = embed_dim self.sample_height = sample_height self.sample_width = sample_width @@ -206,9 +208,15 @@ class CogVideoXPatchEmbed(nn.Module): self.use_positional_embeddings = use_positional_embeddings self.use_learned_positional_embeddings = use_learned_positional_embeddings - self.proj = nn.Conv2d( - in_channels, embed_dim, kernel_size=(patch_size, patch_size), stride=patch_size, bias=bias - ) + if patch_size_t is None: + # CogVideoX 1.0 checkpoints + self.proj = nn.Conv2d( + in_channels, embed_dim, kernel_size=(patch_size, patch_size), stride=patch_size, bias=bias + ) + else: + # CogVideoX 1.5 checkpoints + self.proj = nn.Linear(in_channels * patch_size * patch_size * patch_size_t, embed_dim) + self.text_proj = nn.Linear(text_embed_dim, embed_dim) if use_positional_embeddings or use_learned_positional_embeddings: @@ -247,12 +255,24 @@ class CogVideoXPatchEmbed(nn.Module): """ text_embeds = self.text_proj(text_embeds) - batch, num_frames, channels, height, width = image_embeds.shape - image_embeds = image_embeds.reshape(-1, channels, height, width) - image_embeds = self.proj(image_embeds) - image_embeds = image_embeds.view(batch, num_frames, *image_embeds.shape[1:]) - image_embeds = image_embeds.flatten(3).transpose(2, 3) # [batch, num_frames, height x width, channels] - image_embeds = image_embeds.flatten(1, 2) # [batch, num_frames x height x width, channels] + batch_size, num_frames, channels, height, width = image_embeds.shape + + if self.patch_size_t is None: + image_embeds = image_embeds.reshape(-1, channels, height, width) + image_embeds = self.proj(image_embeds) + image_embeds = image_embeds.view(batch_size, num_frames, *image_embeds.shape[1:]) + image_embeds = image_embeds.flatten(3).transpose(2, 3) # [batch, num_frames, height x width, channels] + image_embeds = image_embeds.flatten(1, 2) # [batch, num_frames x height x width, channels] + else: + p = self.patch_size + p_t = self.patch_size_t + + image_embeds = image_embeds.permute(0, 1, 3, 4, 2) + image_embeds = image_embeds.reshape( + batch_size, num_frames // p_t, p_t, height // p, p, width // p, p, channels + ) + image_embeds = image_embeds.permute(0, 1, 3, 5, 7, 2, 4, 6).flatten(4, 7).flatten(1, 3) + image_embeds = self.proj(image_embeds) embeds = torch.cat( [text_embeds, image_embeds], dim=1 diff --git a/scepter/modules/model/backbone/cogvideox/utils.py b/scepter/modules/model/backbone/cogvideox/utils.py index 03fec8f..23091f3 100644 --- a/scepter/modules/model/backbone/cogvideox/utils.py +++ b/scepter/modules/model/backbone/cogvideox/utils.py @@ -459,7 +459,14 @@ def get_1d_rotary_pos_embed( def get_3d_rotary_pos_embed( - embed_dim, crops_coords, grid_size, temporal_size, theta: int = 10000, use_real: bool = True + embed_dim, + crops_coords, + grid_size, + temporal_size, + theta: int = 10000, + use_real: bool = True, + grid_type: str = "linspace", + max_size: Optional[Tuple[int, int]] = None, ) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]: """ RoPE for video tokens with 3D structure. @@ -475,17 +482,30 @@ def get_3d_rotary_pos_embed( The size of the temporal dimension. theta (`float`): Scaling factor for frequency computation. + grid_type (`str`): + Whether to use "linspace" or "slice" to compute grids. Returns: `torch.Tensor`: positional embedding with shape `(temporal_size * grid_size[0] * grid_size[1], embed_dim/2)`. """ if use_real is not True: raise ValueError(" `use_real = False` is not currently supported for get_3d_rotary_pos_embed") - start, stop = crops_coords - grid_size_h, grid_size_w = grid_size - grid_h = np.linspace(start[0], stop[0], grid_size_h, endpoint=False, dtype=np.float32) - grid_w = np.linspace(start[1], stop[1], grid_size_w, endpoint=False, dtype=np.float32) - grid_t = np.linspace(0, temporal_size, temporal_size, endpoint=False, dtype=np.float32) + + if grid_type == "linspace": + start, stop = crops_coords + grid_size_h, grid_size_w = grid_size + grid_h = np.linspace(start[0], stop[0], grid_size_h, endpoint=False, dtype=np.float32) + grid_w = np.linspace(start[1], stop[1], grid_size_w, endpoint=False, dtype=np.float32) + grid_t = np.arange(temporal_size, dtype=np.float32) + grid_t = np.linspace(0, temporal_size, temporal_size, endpoint=False, dtype=np.float32) + elif grid_type == "slice": + max_h, max_w = max_size + grid_size_h, grid_size_w = grid_size + grid_h = np.arange(max_h, dtype=np.float32) + grid_w = np.arange(max_w, dtype=np.float32) + grid_t = np.arange(temporal_size, dtype=np.float32) + else: + raise ValueError("Invalid value passed for `grid_type`.") # Compute dimensions for each axis dim_t = embed_dim // 4 @@ -521,6 +541,12 @@ def get_3d_rotary_pos_embed( t_cos, t_sin = freqs_t # both t_cos and t_sin has shape: temporal_size, dim_t h_cos, h_sin = freqs_h # both h_cos and h_sin has shape: grid_size_h, dim_h w_cos, w_sin = freqs_w # both w_cos and w_sin has shape: grid_size_w, dim_w + + if grid_type == "slice": + t_cos, t_sin = t_cos[:temporal_size], t_sin[:temporal_size] + h_cos, h_sin = h_cos[:grid_size_h], h_sin[:grid_size_h] + w_cos, w_sin = w_cos[:grid_size_w], w_sin[:grid_size_w] + cos = combine_time_height_width(t_cos, h_cos, w_cos) sin = combine_time_height_width(t_sin, h_sin, w_sin) return cos, sin diff --git a/scepter/modules/model/backbone/flux/__init__.py b/scepter/modules/model/backbone/flux/__init__.py index 81cea17..21cef21 100644 --- a/scepter/modules/model/backbone/flux/__init__.py +++ b/scepter/modules/model/backbone/flux/__init__.py @@ -1,3 +1,3 @@ # -*- coding: utf-8 -*- # Copyright (c) Alibaba, Inc. and its affiliates. -from .flux import Flux +from .flux import Flux, FluxMR, FluxMRFill, FluxMRRedux, FluxMRControl diff --git a/scepter/modules/model/backbone/flux/flux.py b/scepter/modules/model/backbone/flux/flux.py index 97cbf94..56919d4 100644 --- a/scepter/modules/model/backbone/flux/flux.py +++ b/scepter/modules/model/backbone/flux/flux.py @@ -1,6 +1,9 @@ # -*- coding: utf-8 -*- # Copyright (c) Alibaba, Inc. and its affiliates. +# This file contains code that is adapted from +# https://github.com/black-forest-labs/flux.git import math +from collections import OrderedDict from functools import partial import torch @@ -15,8 +18,6 @@ from torch.utils.checkpoint import checkpoint_sequential from torch.nn.utils.rnn import pad_sequence from .layers import (DoubleStreamBlock, EmbedND, LastLayer, MLPEmbedder, SingleStreamBlock, timestep_embedding) - - @BACKBONES.register_class() class Flux(BaseModel): """ @@ -98,7 +99,14 @@ class Flux(BaseModel): qkv_bias = cfg.QKV_BIAS depth = cfg.DEPTH depth_single_blocks = cfg.DEPTH_SINGLE_BLOCKS - self.use_grad_checkpoint = cfg.get('USE_GRAD_CHECKPOINT', False) + self.use_grad_checkpoint = cfg.get("USE_GRAD_CHECKPOINT", False) + self.attn_backend = cfg.get("ATTN_BACKEND", "pytorch") + self.cache_pretrain_model = cfg.get("CACHE_PRETRAIN_MODEL", False) + self.lora_model = cfg.get("DIFFUSERS_LORA_MODEL", None) + self.comfyui_lora_model = cfg.get("COMFYUI_LORA_MODEL", None) + self.swift_lora_model = cfg.get("SWIFT_LORA_MODEL", None) + self.blackforest_lora_model = cfg.get("BLACKFOREST_LORA_MODEL", None) + self.pretrain_adapter = cfg.get("PRETRAIN_ADAPTER", None) if hidden_size % num_heads != 0: raise ValueError( @@ -119,85 +127,350 @@ class Flux(BaseModel): if self.guidance_embed else nn.Identity()) self.txt_in = nn.Linear(context_in_dim, self.hidden_size) - self.double_blocks = nn.ModuleList([ - DoubleStreamBlock( - self.hidden_size, - self.num_heads, - mlp_ratio=mlp_ratio, - qkv_bias=qkv_bias, - ) for _ in range(depth) - ]) + self.double_blocks = nn.ModuleList( + [ + DoubleStreamBlock( + self.hidden_size, + self.num_heads, + mlp_ratio=mlp_ratio, + qkv_bias=qkv_bias, + backend=self.attn_backend + ) + for _ in range(depth) + ] + ) - self.single_blocks = nn.ModuleList([ - SingleStreamBlock(self.hidden_size, - self.num_heads, - mlp_ratio=mlp_ratio) - for _ in range(depth_single_blocks) - ]) + self.single_blocks = nn.ModuleList( + [ + SingleStreamBlock(self.hidden_size, self.num_heads, mlp_ratio=mlp_ratio, backend=self.attn_backend) + for _ in range(depth_single_blocks) + ] + ) self.final_layer = LastLayer(self.hidden_size, 1, self.out_channels) def prepare_input(self, x, context, y, x_shape=None): # x.shape [6, 16, 16, 16] target is [6, 16, 768, 1360] bs, c, h, w = x.shape - x = rearrange(x, 'b c (h ph) (w pw) -> b (h w) (c ph pw)', ph=2, pw=2) + x = rearrange(x, "b c (h ph) (w pw) -> b (h w) (c ph pw)", ph=2, pw=2) x_id = torch.zeros(h // 2, w // 2, 3) x_id[..., 1] = x_id[..., 1] + torch.arange(h // 2)[:, None] x_id[..., 2] = x_id[..., 2] + torch.arange(w // 2)[None, :] - x_ids = repeat(x_id, 'h w c -> b (h w) c', b=bs) + x_ids = repeat(x_id, "h w c -> b (h w) c", b=bs) txt_ids = torch.zeros(bs, context.shape[1], 3) return x, x_ids.to(x), context.to(x), txt_ids.to(x), y.to(x), h, w def unpack(self, x: Tensor, height: int, width: int) -> Tensor: return rearrange( x, - 'b (h w) (c ph pw) -> b c (h ph) (w pw)', - h=math.ceil(height / 2), - w=math.ceil(width / 2), + "b (h w) (c ph pw) -> b c (h ph) (w pw)", + h=math.ceil(height/2), + w=math.ceil(width/2), ph=2, pw=2, ) - def load_pretrained_model(self, pretrained_model): - if next(self.parameters()).device.type == 'meta': - map_location = we.device_id - else: - map_location = 'cpu' - if pretrained_model is not None: - with FS.get_from(pretrained_model, - wait_finish=True) as local_model: - if local_model.endswith('safetensors'): - from safetensors.torch import load_file as load_safetensors - sd = load_safetensors(local_model, device=map_location) + def merge_diffuser_lora(self, ori_sd, lora_sd, scale=1.0): + key_map = { + "single_blocks.{}.linear1.weight": {"key_list": [ + ["transformer.single_transformer_blocks.{}.attn.to_q.lora_A.weight", + "transformer.single_transformer_blocks.{}.attn.to_q.lora_B.weight", [0, 3072]], + ["transformer.single_transformer_blocks.{}.attn.to_k.lora_A.weight", + "transformer.single_transformer_blocks.{}.attn.to_k.lora_B.weight", [3072, 6144]], + ["transformer.single_transformer_blocks.{}.attn.to_v.lora_A.weight", + "transformer.single_transformer_blocks.{}.attn.to_v.lora_B.weight", [6144, 9216]], + ["transformer.single_transformer_blocks.{}.proj_mlp.lora_A.weight", + "transformer.single_transformer_blocks.{}.proj_mlp.lora_B.weight", [9216, 21504]] + ], "num": 38}, + "single_blocks.{}.modulation.lin.weight": {"key_list": [ + ["transformer.single_transformer_blocks.{}.norm.linear.lora_A.weight", + "transformer.single_transformer_blocks.{}.norm.linear.lora_B.weight", [0, 9216]], + ], "num": 38}, + "single_blocks.{}.linear2.weight": {"key_list": [ + ["transformer.single_transformer_blocks.{}.proj_out.lora_A.weight", + "transformer.single_transformer_blocks.{}.proj_out.lora_B.weight", [0, 3072]], + ], "num": 38}, + "double_blocks.{}.txt_attn.qkv.weight": {"key_list": [ + ["transformer.transformer_blocks.{}.attn.add_q_proj.lora_A.weight", + "transformer.transformer_blocks.{}.attn.add_q_proj.lora_B.weight", [0, 3072]], + ["transformer.transformer_blocks.{}.attn.add_k_proj.lora_A.weight", + "transformer.transformer_blocks.{}.attn.add_k_proj.lora_B.weight", [3072, 6144]], + ["transformer.transformer_blocks.{}.attn.add_v_proj.lora_A.weight", + "transformer.transformer_blocks.{}.attn.add_v_proj.lora_B.weight", [6144, 9216]], + ], "num": 19}, + "double_blocks.{}.img_attn.qkv.weight": {"key_list": [ + ["transformer.transformer_blocks.{}.attn.to_q.lora_A.weight", + "transformer.transformer_blocks.{}.attn.to_q.lora_B.weight", [0, 3072]], + ["transformer.transformer_blocks.{}.attn.to_k.lora_A.weight", + "transformer.transformer_blocks.{}.attn.to_k.lora_B.weight", [3072, 6144]], + ["transformer.transformer_blocks.{}.attn.to_v.lora_A.weight", + "transformer.transformer_blocks.{}.attn.to_v.lora_B.weight", [6144, 9216]], + ], "num": 19}, + "double_blocks.{}.img_attn.proj.weight": {"key_list": [ + ["transformer.transformer_blocks.{}.attn.to_out.0.lora_A.weight", + "transformer.transformer_blocks.{}.attn.to_out.0.lora_B.weight", [0, 3072]] + ], "num": 19}, + "double_blocks.{}.txt_attn.proj.weight": {"key_list": [ + ["transformer.transformer_blocks.{}.attn.to_add_out.lora_A.weight", + "transformer.transformer_blocks.{}.attn.to_add_out.lora_B.weight", [0, 3072]] + ], "num": 19}, + "double_blocks.{}.img_mlp.0.weight": {"key_list": [ + ["transformer.transformer_blocks.{}.ff.net.0.proj.lora_A.weight", + "transformer.transformer_blocks.{}.ff.net.0.proj.lora_B.weight", [0, 12288]] + ], "num": 19}, + "double_blocks.{}.img_mlp.2.weight": {"key_list": [ + ["transformer.transformer_blocks.{}.ff.net.2.lora_A.weight", + "transformer.transformer_blocks.{}.ff.net.2.lora_B.weight", [0, 3072]] + ], "num": 19}, + "double_blocks.{}.txt_mlp.0.weight": {"key_list": [ + ["transformer.transformer_blocks.{}.ff_context.net.0.proj.lora_A.weight", + "transformer.transformer_blocks.{}.ff_context.net.0.proj.lora_B.weight", [0, 12288]] + ], "num": 19}, + "double_blocks.{}.txt_mlp.2.weight": {"key_list": [ + ["transformer.transformer_blocks.{}.ff_context.net.2.lora_A.weight", + "transformer.transformer_blocks.{}.ff_context.net.2.lora_B.weight", [0, 3072]] + ], "num": 19}, + "double_blocks.{}.img_mod.lin.weight": {"key_list": [ + ["transformer.transformer_blocks.{}.norm1.linear.lora_A.weight", + "transformer.transformer_blocks.{}.norm1.linear.lora_B.weight", [0, 18432]] + ], "num": 19}, + "double_blocks.{}.txt_mod.lin.weight": {"key_list": [ + ["transformer.transformer_blocks.{}.norm1_context.linear.lora_A.weight", + "transformer.transformer_blocks.{}.norm1_context.linear.lora_B.weight", [0, 18432]] + ], "num": 19} + } + cover_lora_keys = set() + cover_ori_keys = set() + for k, v in key_map.items(): + key_list = v["key_list"] + block_num = v["num"] + for block_id in range(block_num): + for k_list in key_list: + if k_list[0].format(block_id) in lora_sd and k_list[1].format(block_id) in lora_sd: + cover_lora_keys.add(k_list[0].format(block_id)) + cover_lora_keys.add(k_list[1].format(block_id)) + current_weight = torch.matmul(lora_sd[k_list[0].format(block_id)].permute(1, 0), + lora_sd[k_list[1].format(block_id)].permute(1, 0)).permute(1, 0) + ori_sd[k.format(block_id)][k_list[2][0]:k_list[2][1], ...] += scale * current_weight + cover_ori_keys.add(k.format(block_id)) + # lora_sd.pop(k_list[0].format(block_id)) + # lora_sd.pop(k_list[1].format(block_id)) + self.logger.info(f"merge_blackforest_lora loads lora'parameters lora-paras: \n" + f"cover-{len(cover_lora_keys)} vs total {len(lora_sd)} \n" + f"cover ori-{len(cover_ori_keys)} vs total {len(ori_sd)}") + return ori_sd + + def merge_swift_lora(self, ori_sd, lora_sd, scale = 1.0): + have_lora_keys = {} + for k, v in lora_sd.items(): + k = k[len("model."):] if k.startswith("model.") else k + ori_key = k.split("lora")[0] + "weight" + if ori_key not in ori_sd: + raise f"{ori_key} should in the original statedict" + if ori_key not in have_lora_keys: + have_lora_keys[ori_key] = {} + if "lora_A" in k: + have_lora_keys[ori_key]["lora_A"] = v + elif "lora_B" in k: + have_lora_keys[ori_key]["lora_B"] = v + else: + raise NotImplementedError + self.logger.info(f"merge_swift_lora loads lora'parameters {len(have_lora_keys)}") + for key, v in have_lora_keys.items(): + current_weight = torch.matmul(v["lora_A"].permute(1, 0), v["lora_B"].permute(1, 0)).permute(1, 0) + ori_sd[key] += scale * current_weight + return ori_sd + + + def merge_blackforest_lora(self, ori_sd, lora_sd, scale = 1.0): + have_lora_keys = {} + cover_lora_keys = set() + cover_ori_keys = set() + for k, v in lora_sd.items(): + if "lora" in k: + ori_key = k.split("lora")[0] + "weight" + if ori_key not in ori_sd: + raise f"{ori_key} should in the original statedict" + if ori_key not in have_lora_keys: + have_lora_keys[ori_key] = {} + if "lora_A" in k: + have_lora_keys[ori_key]["lora_A"] = v + cover_lora_keys.add(k) + cover_ori_keys.add(ori_key) + elif "lora_B" in k: + have_lora_keys[ori_key]["lora_B"] = v + cover_lora_keys.add(k) + cover_ori_keys.add(ori_key) + else: + if k in ori_sd: + ori_sd[k] = v + cover_lora_keys.add(k) + cover_ori_keys.add(k) else: - sd = torch.load(local_model, map_location=map_location) - missing, unexpected = self.load_state_dict(sd, - strict=False, - assign=True) + print("unsurpport keys: ", k) + self.logger.info(f"merge_blackforest_lora loads lora'parameters lora-paras: \n" + f"cover-{len(cover_lora_keys)} vs total {len(lora_sd)} \n" + f"cover ori-{len(cover_ori_keys)} vs total {len(ori_sd)}") + + for key, v in have_lora_keys.items(): + current_weight = torch.matmul(v["lora_A"].permute(1, 0), v["lora_B"].permute(1, 0)).permute(1, 0) + # print(key, ori_sd[key].shape, current_weight.shape) + ori_sd[key] += scale * current_weight + return ori_sd + + def merge_comfyui_lora(self, ori_sd, lora_sd, scale = 1.0): + ori_key_map = {key.replace("_", ".") : key for key in ori_sd.keys()} + parse_ckpt = OrderedDict() + for k, v in lora_sd.items(): + if "alpha" in k: + continue + k = k.replace("lora_unet_", "").replace("_", ".") + map_k = ori_key_map[k.split(".lora")[0] + ".weight"] + if map_k not in parse_ckpt: + parse_ckpt[map_k] = {} + if "lora.up" in k: + parse_ckpt[map_k]["lora_up"] = v + elif "lora.down" in k: + parse_ckpt[map_k]["lora_down"] = v + if self.cache_pretrain_model: + self.lora_dict[self.comfyui_lora_model] = {} + + for key, v in parse_ckpt.items(): + current_weight = torch.matmul(v["lora_down"].permute(1, 0), v["lora_up"].permute(1, 0)).permute(1, 0) + self.lora_dict[self.comfyui_lora_model] = current_weight + ori_sd[key] += scale * current_weight + return ori_sd + + def easy_lora_merge(self, ori_sd, lora_sd, scale = 1.0): + for key, v in lora_sd.items(): + ori_sd[key] += scale * v + return ori_sd + + def load_pretrained_model(self, pretrained_model, lora_scale = 1.0): + if next(self.parameters()).device.type == 'meta': + map_location = torch.device(we.device_id) + safe_device = we.device_id + else: + map_location = "cpu" + safe_device = "cpu" + + if pretrained_model is not None: + if not hasattr(self, "ckpt"): + with FS.get_from(pretrained_model, wait_finish=True) as local_model: + if local_model.endswith('safetensors'): + from safetensors.torch import load_file as load_safetensors + ckpt = load_safetensors(local_model, device=safe_device) + else: + ckpt = torch.load(local_model, map_location=map_location, weights_only=True) + if "state_dict" in ckpt: + ckpt = ckpt["state_dict"] + if "model" in ckpt: + ckpt = ckpt["model"]["model"] + if self.cache_pretrain_model: + self.ckpt = ckpt + self.lora_dict = {} + else: + ckpt = self.ckpt + + new_ckpt = OrderedDict() + for k, v in ckpt.items(): + if k in ("img_in.weight"): + model_p = self.state_dict()[k] + if v.shape != model_p.shape: + expanded_state_dict_weight = torch.zeros_like(model_p, device=v.device) + slices = tuple(slice(0, dim) for dim in v.shape) + expanded_state_dict_weight[slices] = v + new_ckpt[k] = expanded_state_dict_weight + else: + new_ckpt[k] = v + else: + new_ckpt[k] = v + + + if self.lora_model is not None: + with FS.get_from(self.lora_model, wait_finish=True) as local_model: + if local_model.endswith('safetensors'): + from safetensors.torch import load_file as load_safetensors + lora_sd = load_safetensors(local_model, device=safe_device) + else: + lora_sd = torch.load(local_model, map_location=map_location, weights_only=True) + new_ckpt = self.merge_diffuser_lora(new_ckpt, lora_sd, scale=lora_scale) + if self.swift_lora_model is not None: + if not isinstance(self.swift_lora_model, list): + self.swift_lora_model = [(self.swift_lora_model, 1.0)] + for lora_model in self.swift_lora_model: + if isinstance(lora_model, str): + lora_model = (lora_model, 1.0/len(self.swift_lora_model)) + print(lora_model) + self.logger.info(f"load swift lora model: {lora_model}") + with FS.get_from(lora_model[0], wait_finish=True) as local_model: + if local_model.endswith('safetensors'): + from safetensors.torch import load_file as load_safetensors + lora_sd = load_safetensors(local_model, device=safe_device) + else: + lora_sd = torch.load(local_model, map_location=map_location, weights_only=True) + new_ckpt = self.merge_swift_lora(new_ckpt, lora_sd, scale=lora_model[1]) + + if self.blackforest_lora_model is not None: + with FS.get_from(self.blackforest_lora_model, wait_finish=True) as local_model: + if local_model.endswith('safetensors'): + from safetensors.torch import load_file as load_safetensors + lora_sd = load_safetensors(local_model, device=safe_device) + else: + lora_sd = torch.load(local_model, map_location=map_location, weights_only=True) + new_ckpt = self.merge_blackforest_lora(new_ckpt, lora_sd, scale=lora_scale) + + if self.comfyui_lora_model is not None: + if hasattr(self, "current_lora") and self.current_lora == self.comfyui_lora_model: + return + if hasattr(self, "lora_dict") and self.comfyui_lora_model in self.lora_dict: + new_ckpt = self.easy_lora_merge(new_ckpt, self.lora_dict[self.comfyui_lora_model], scale=lora_scale) + else: + with FS.get_from(self.comfyui_lora_model, wait_finish=True) as local_model: + if local_model.endswith('safetensors'): + from safetensors.torch import load_file as load_safetensors + lora_sd = load_safetensors(local_model, device=safe_device) + else: + lora_sd = torch.load(local_model, map_location=map_location, weights_only=True) + new_ckpt = self.merge_comfyui_lora(new_ckpt, lora_sd, scale=lora_scale) + if self.comfyui_lora_model: + self.current_lora = self.comfyui_lora_model + + + adapter_ckpt = {} + if self.pretrain_adapter is not None: + with FS.get_from(self.pretrain_adapter, wait_finish=True) as local_adapter: + if local_adapter.endswith('safetensors'): + from safetensors.torch import load_file as load_safetensors + adapter_ckpt = load_safetensors(local_adapter, device=safe_device) + else: + adapter_ckpt = torch.load(local_adapter, map_location=map_location, weights_only=True) + new_ckpt.update(adapter_ckpt) + + missing, unexpected = self.load_state_dict(new_ckpt, strict=False, assign=True) self.logger.info( f'Restored from {pretrained_model} with {len(missing)} missing and {len(unexpected)} unexpected keys' ) if len(missing) > 0: - self.logger.info(f'Missing Keys:\n {missing}') # noqa + self.logger.info(f'Missing Keys:\n {missing}') if len(unexpected) > 0: - self.logger.info(f'\nUnexpected Keys:\n {unexpected}') # noqa + self.logger.info(f'\nUnexpected Keys:\n {unexpected}') - def forward(self, - x: Tensor, - t: Tensor, - cond: dict = {}, - guidance: Tensor | None = None, - gc_seg: int = 0) -> Tensor: - x, x_ids, txt, txt_ids, y, h, w = self.prepare_input( - x, cond['context'], cond['y']) + def forward( + self, + x: Tensor, + t: Tensor, + cond: dict = {}, + guidance: Tensor | None = None, + gc_seg: int = 0 + ) -> Tensor: + x, x_ids, txt, txt_ids, y, h, w = self.prepare_input(x, cond["context"], cond["y"]) # running on sequences img x = self.img_in(x) vec = self.time_in(timestep_embedding(t, 256)) if self.guidance_embed: if guidance is None: - raise ValueError( - "Didn't get guidance strength for guidance distilled model." - ) + raise ValueError("Didn't get guidance strength for guidance distilled model.") vec = vec + self.guidance_in(timestep_embedding(guidance, 256)) vec = vec + self.vector_in(y) txt = self.txt_in(txt) @@ -211,12 +484,11 @@ class Flux(BaseModel): x = torch.cat((txt, x), 1) if self.use_grad_checkpoint and gc_seg >= 0: x = checkpoint_sequential( - functions=[ - partial(block, **kwargs) for block in self.double_blocks - ], + functions=[partial(block, **kwargs) for block in self.double_blocks], segments=gc_seg if gc_seg > 0 else len(self.double_blocks), input=x, - use_reentrant=False) + use_reentrant=False + ) else: for block in self.double_blocks: x = block(x, **kwargs) @@ -228,18 +500,16 @@ class Flux(BaseModel): if self.use_grad_checkpoint and gc_seg >= 0: x = checkpoint_sequential( - functions=[ - partial(block, **kwargs) for block in self.single_blocks - ], + functions=[partial(block, **kwargs) for block in self.single_blocks], segments=gc_seg if gc_seg > 0 else len(self.single_blocks), input=x, - use_reentrant=False) + use_reentrant=False + ) else: for block in self.single_blocks: x = block(x, **kwargs) - x = x[:, txt.shape[1]:, ...] - x = self.final_layer( - x, vec) # (N, T, patch_size ** 2 * out_channels) 6 64 64 + x = x[:, txt.shape[1] :, ...] + x = self.final_layer(x, vec) # (N, T, patch_size ** 2 * out_channels) 6 64 64 x = self.unpack(x, h, w) return x @@ -249,11 +519,13 @@ class Flux(BaseModel): __class__.__name__, Flux.para_dict, set_name=True) - @BACKBONES.register_class() class FluxMR(Flux): def prepare_input(self, x, cond): - context, y = cond["context"].to(x), cond["y"].to(x) + if isinstance(cond['context'], list): + context, y = torch.cat(cond["context"], dim=0).to(x), torch.cat(cond["y"], dim=0).to(x) + else: + context, y = cond['context'].to(x), cond['y'].to(x) batch_frames, batch_frames_ids = [], [] for ix, shape in zip(x, cond["x_shapes"]): # unpack image from sequence @@ -319,7 +591,7 @@ class FluxMR(Flux): x, x_ids, txt, txt_ids, y, mask_x, mask_txt, seq_length_list = self.prepare_input(x, cond) # running on sequences img vec = self.time_in(timestep_embedding(t, 256)) - if self.guidance_embed: + if self.guidance_embed and guidance[-1] >= 0: if guidance is None: raise ValueError("Didn't get guidance strength for guidance distilled model.") vec = vec + self.guidance_in(timestep_embedding(guidance, 256)) @@ -371,7 +643,170 @@ class FluxMR(Flux): @staticmethod def get_config_template(): - return dict_to_yaml('BACKBONE', + return dict_to_yaml('MODEL', __class__.__name__, FluxMR.para_dict, set_name=True) +@BACKBONES.register_class() +class FluxMRFill(FluxMR): + def __init__(self, cfg, logger = None): + super().__init__(cfg, logger) + def prepare_input(self, x, cond): + context, y = cond["context"], cond["y"] + batch_frames, batch_frames_ids = [], [] + for ix, shape, imask, ie, ie_mask in zip(x, cond["x_shapes"], cond["x_mask"], + cond["edit"], cond["edit_mask"]): + # unpack image from sequence + ix = ix[:, :shape[0] * shape[1]].view(-1, shape[0], shape[1]) + imask = torch.ones_like(ix[[0], :, :]) if imask is None else imask.squeeze(0) + if len(ie) > 0: + ie = ie[0].squeeze(0) + ie_mask = torch.ones((ix.shape[0] * 4, ix.shape[1], ix.shape[2])) if ie_mask is None else ie_mask[0].squeeze(0) + else: + ie, ie_mask = torch.zeros_like(ix).to(x), torch.ones_like(imask).to(x) + ix = torch.cat([ix, ie, ie_mask], dim=0) + c, h, w = ix.shape + ix = rearrange(ix, "c (h ph) (w pw) -> (h w) (c ph pw)", ph=2, pw=2) + ix_id = torch.zeros(h // 2, w // 2, 3) + ix_id[..., 1] = ix_id[..., 1] + torch.arange(h // 2)[:, None] + ix_id[..., 2] = ix_id[..., 2] + torch.arange(w // 2)[None, :] + ix_id = rearrange(ix_id, "h w c -> (h w) c") + batch_frames.append([ix]) + batch_frames_ids.append([ix_id]) + x_list, x_id_list, mask_x_list, x_seq_length = [], [], [], [] + for frames, frame_ids in zip(batch_frames, batch_frames_ids): + proj_frames = [] + for idx, one_frame in enumerate(frames): + one_frame = self.img_in(one_frame) + proj_frames.append(one_frame) + ix = torch.cat(proj_frames, dim=0) + if_id = torch.cat(frame_ids, dim=0) + x_list.append(ix) + x_id_list.append(if_id) + mask_x_list.append(torch.ones(ix.shape[0]).to(ix.device, non_blocking=True).bool()) + x_seq_length.append(ix.shape[0]) + # if len(x_list) < 1: import pdb;pdb.set_trace() + x = pad_sequence(tuple(x_list), batch_first=True) + x_ids = pad_sequence(tuple(x_id_list), batch_first=True).to(x) # [b,pad_seq,2] pad (0.,0.) at dim2 + mask_x = pad_sequence(tuple(mask_x_list), batch_first=True) + # import pdb;pdb.set_trace() + if isinstance(context, list): + txt_list, mask_txt_list, y_list = [], [], [] + for sample_id, (ctx, yy) in enumerate(zip(context, y)): + txt_list.append(self.txt_in(ctx.to(x))) + mask_txt_list.append(torch.ones(txt_list[-1].shape[0]).to(ctx.device, non_blocking=True).bool()) + y_list.append(yy.to(x)) + txt = pad_sequence(tuple(txt_list), batch_first=True) + txt_ids = torch.zeros(txt.shape[0], txt.shape[1], 3).to(x) + mask_txt = pad_sequence(tuple(mask_txt_list), batch_first=True) + y = torch.cat(y_list, dim=0) + assert y.ndim == 2 and txt.ndim == 3 + else: + txt = self.txt_in(context) + txt_ids = torch.zeros(context.shape[0], context.shape[1], 3).to(x) + mask_txt = torch.ones(context.shape[0], context.shape[1]).to(x.device, non_blocking=True).bool() + return x, x_ids, txt, txt_ids, y, mask_x, mask_txt, x_seq_length + + @staticmethod + def get_config_template(): + return dict_to_yaml('MODEL', + __class__.__name__, + FluxMRFill.para_dict, + set_name=True) +@BACKBONES.register_class() +class FluxMRRedux(FluxMR): + ''' + ref_image_siglip + projector + ''' + def __init__(self, cfg, logger = None): + super().__init__(cfg, logger) + self.redux_dim = cfg.get("REDUX_DIM", 1152) + self.context_in_dim = cfg.CONTEXT_IN_DIM + self.redux_up = nn.Linear(self.redux_dim, self.context_in_dim * 3) + self.redux_down = nn.Linear(self.context_in_dim * 3, self.context_in_dim) + + + def prepare_input(self, x, cond): + ref_x = cond.get("ref_x", None) + context, y = torch.cat(cond["context"], dim=0).to(x), torch.cat(cond["y"], dim=0).to(x) + if ref_x is not None: + ref_x = [torch.cat(ref_ix, dim=0).mean(dim=0, keepdim=True) for ref_ix in ref_x] + ref_x = self.redux_down(nn.functional.silu(self.redux_up(torch.cat(ref_x, dim=0)))) + context = torch.cat((context, ref_x), dim=-2) + + batch_frames, batch_frames_ids = [], [] + for ix, shape in zip(x, cond["x_shapes"]): + # unpack image from sequence + ix = ix[:, :shape[0] * shape[1]].view(-1, shape[0], shape[1]) + c, h, w = ix.shape + ix = rearrange(ix, "c (h ph) (w pw) -> (h w) (c ph pw)", ph=2, pw=2) + ix_id = torch.zeros(h // 2, w // 2, 3) + ix_id[..., 1] = ix_id[..., 1] + torch.arange(h // 2)[:, None] + ix_id[..., 2] = ix_id[..., 2] + torch.arange(w // 2)[None, :] + ix_id = rearrange(ix_id, "h w c -> (h w) c") + batch_frames.append([ix]) + batch_frames_ids.append([ix_id]) + + x_list, x_id_list, mask_x_list, x_seq_length = [], [], [], [] + for frames, frame_ids in zip(batch_frames, batch_frames_ids): + proj_frames = [] + for idx, one_frame in enumerate(frames): + one_frame = self.img_in(one_frame) + proj_frames.append(one_frame) + ix = torch.cat(proj_frames, dim=0) + if_id = torch.cat(frame_ids, dim=0) + x_list.append(ix) + x_id_list.append(if_id) + mask_x_list.append(torch.ones(ix.shape[0]).to(ix.device, non_blocking=True).bool()) + x_seq_length.append(ix.shape[0]) + x = pad_sequence(tuple(x_list), batch_first=True) + x_ids = pad_sequence(tuple(x_id_list), batch_first=True).to(x) # [b,pad_seq,2] pad (0.,0.) at dim2 + mask_x = pad_sequence(tuple(mask_x_list), batch_first=True) + + txt = self.txt_in(context) + txt_ids = torch.zeros(context.shape[0], context.shape[1], 3).to(x) + mask_txt = torch.ones(context.shape[0], context.shape[1]).to(x.device, non_blocking=True).bool() + + return x, x_ids, txt, txt_ids, y, mask_x, mask_txt, x_seq_length + @staticmethod + def get_config_template(): + return dict_to_yaml('MODEL', + __class__.__name__, + FluxMRRedux.para_dict, + set_name=True) +@BACKBONES.register_class() +class FluxMRControl(FluxMR): + ''' + cat([x, ie]) ensure the same size bettwn the x and ie + ''' + def prepare_input(self, x, cond, *args, **kwargs ): + context, y = torch.cat(cond["context"], dim=0).to(x), torch.cat(cond["y"], dim=0).to(x) + x_list, x_id_list, mask_x_list, x_seq_length = [], [], [], [] + for ix, shape, ie in zip(x, cond["x_shapes"], cond["edit"]): + ix = ix[:, :shape[0] * shape[1]].view(-1, shape[0], shape[1]) + ix = torch.cat([ix, ie], dim=0) + c, h, w = ix.shape + ix = rearrange(ix, "c (h ph) (w pw) -> (h w) (c ph pw)", ph=2, pw=2) + ix_id = torch.zeros(h // 2, w // 2, 3) + ix_id[..., 1] = ix_id[..., 1] + torch.arange(h // 2)[:, None] + ix_id[..., 2] = ix_id[..., 2] + torch.arange(w // 2)[None, :] + ix_id = rearrange(ix_id, "h w c -> (h w) c") + x_list.append(self.img_in(ix)) + x_id_list.append(ix_id) + mask_x_list.append(torch.ones(ix.shape[0]).to(ix.device, non_blocking=True).bool()) + x_seq_length.append(ix.shape[0]) + # if len(x_list) < 1: import pdb;pdb.set_trace() + x = pad_sequence(tuple(x_list), batch_first=True) + x_ids = pad_sequence(tuple(x_id_list), batch_first=True).to(x) # [b,pad_seq,2] pad (0.,0.) at dim2 + mask_x = pad_sequence(tuple(mask_x_list), batch_first=True) + txt = self.txt_in(context) + txt_ids = torch.zeros(context.shape[0], context.shape[1], 3).to(x) + mask_txt = torch.ones(context.shape[0], context.shape[1]).to(x.device, non_blocking=True).bool() + return x, x_ids, txt, txt_ids, y, mask_x, mask_txt, x_seq_length + + @staticmethod + def get_config_template(): + return dict_to_yaml('MODEL', + __class__.__name__, + FluxMRControl.para_dict, + set_name=True) diff --git a/scepter/modules/model/backbone/flux/layers.py b/scepter/modules/model/backbone/flux/layers.py index 9a855d3..044be96 100644 --- a/scepter/modules/model/backbone/flux/layers.py +++ b/scepter/modules/model/backbone/flux/layers.py @@ -1,5 +1,7 @@ # -*- coding: utf-8 -*- # Copyright (c) Alibaba, Inc. and its affiliates. +# This file contains code that is adapted from +# https://github.com/black-forest-labs/flux.git from __future__ import annotations import math @@ -351,7 +353,7 @@ class SingleStreamBlock(nn.Module): if mask is not None: mask = repeat(mask, 'B L S-> B H L S', H=self.num_heads) # compute attention - attn = attention(q, k, v, pe=pe, mask=mask) + attn = attention(q, k, v, pe=pe, mask=mask, backend=self.backend) # compute activation in mlp stream, cat again and run second linear layer output = self.linear2(torch.cat((attn, self.mlp_act(mlp)), 2)) return x + mod.gate * output diff --git a/scepter/modules/model/backbone/image/vit_modify.py b/scepter/modules/model/backbone/image/vit_modify.py index 7a3e72f..5366582 100644 --- a/scepter/modules/model/backbone/image/vit_modify.py +++ b/scepter/modules/model/backbone/image/vit_modify.py @@ -74,7 +74,7 @@ class VisualTransformer(BaseModel): with FS.get_from(self.pretrain_path, wait_finish=True) as local_file: logger.info(f'Loading checkpoint from {self.pretrain_path}') - visual_pre = torch.load(local_file, map_location='cpu') + visual_pre = torch.load(local_file, map_location='cpu', weights_only=True) if not use_proj: visual_pre.pop('proj') if visual_pre['conv1.weight'].dtype == torch.float16: @@ -145,7 +145,7 @@ class SomeFTVisualTransformer(BaseModel): with FS.get_from(self.pretrain_path, wait_finish=True) as local_file: logger.info(f'Loading checkpoint from {self.pretrain_path}') - visual_pre = torch.load(local_file, map_location='cpu') + visual_pre = torch.load(local_file, map_location='cpu', weights_only=True) state_dict_update = self.reformat_state_dict(visual_pre) self.visual.load_state_dict(state_dict_update, strict=True) diff --git a/scepter/modules/model/backbone/mmdit/sd3.py b/scepter/modules/model/backbone/mmdit/sd3.py index 43161b5..caf6916 100644 --- a/scepter/modules/model/backbone/mmdit/sd3.py +++ b/scepter/modules/model/backbone/mmdit/sd3.py @@ -1136,7 +1136,7 @@ class MMDiT(BaseModel): from safetensors.torch import load_file as load_safetensors model = load_safetensors(local_path) else: - model = torch.load(local_path, map_location='cpu') + model = torch.load(local_path, map_location='cpu', weights_only=True) if 'state_dict' in model: model = model['state_dict'] new_ckpt = OrderedDict() diff --git a/scepter/modules/model/backbone/pixart/pixart_alpha.py b/scepter/modules/model/backbone/pixart/pixart_alpha.py index 7febcbb..4914f3d 100644 --- a/scepter/modules/model/backbone/pixart/pixart_alpha.py +++ b/scepter/modules/model/backbone/pixart/pixart_alpha.py @@ -354,7 +354,7 @@ class PixArt(BaseModel): 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') + model = torch.load(local_path, map_location='cpu', weights_only=True) if 'state_dict' in model: model = model['state_dict'] new_ckpt = OrderedDict() diff --git a/scepter/modules/model/backbone/transformer/__init__.py b/scepter/modules/model/backbone/transformer/__init__.py index e69de29..cc26a06 100644 --- a/scepter/modules/model/backbone/transformer/__init__.py +++ b/scepter/modules/model/backbone/transformer/__init__.py @@ -0,0 +1,2 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. diff --git a/scepter/modules/model/backbone/transformer/attention.py b/scepter/modules/model/backbone/transformer/attention.py index 15208d9..ea22cde 100644 --- a/scepter/modules/model/backbone/transformer/attention.py +++ b/scepter/modules/model/backbone/transformer/attention.py @@ -10,7 +10,7 @@ import warnings import torch import torch.nn as nn -from torch.cuda import amp +from torch import amp from torch.nn import functional as F from torch.nn.utils.rnn import pad_sequence from tqdm import tqdm @@ -440,7 +440,7 @@ def multi_head_varlen_attention(q_img, k = k.type(flash_dtype) v = v.type(flash_dtype) - with amp.autocast(): + with amp.autocast("cuda"): x = flash_attn_varlen_func(q=q, k=k, v=v, diff --git a/scepter/modules/model/backbone/transformer/pos_embed.py b/scepter/modules/model/backbone/transformer/pos_embed.py index 5299cf6..fa69b15 100644 --- a/scepter/modules/model/backbone/transformer/pos_embed.py +++ b/scepter/modules/model/backbone/transformer/pos_embed.py @@ -13,7 +13,7 @@ import torch.nn as nn import torch.nn.functional as F from einops import rearrange from torch import Tensor -from torch.cuda import amp +from torch import amp from torch.nn.utils.rnn import pad_sequence @@ -175,7 +175,7 @@ def frame_unpad(x, shapes): return torch.concat(frames) -@amp.autocast(enabled=False) +@amp.autocast("cuda", enabled=False) def rope_params(max_seq_len, dim, theta=10000): """ Precompute the frequency tensor for complex exponentials. @@ -189,7 +189,7 @@ def rope_params(max_seq_len, dim, theta=10000): return freqs -@amp.autocast(enabled=False) +@amp.autocast("cuda", enabled=False) def rope_apply(x, grid_sizes, freqs): """ x: [B, L, N, C]. @@ -225,7 +225,7 @@ def rope_apply(x, grid_sizes, freqs): return torch.stack(output) -@amp.autocast(enabled=False) +@amp.autocast("cuda", enabled=False) def rope_apply_multires_pad(x, x_lens, x_shapes, freqs, pad=True): """ x: [B, L, N, C]. @@ -267,7 +267,7 @@ def rope_apply_multires_pad(x, x_lens, x_shapes, freqs, pad=True): return torch.stack(output) if pad else torch.concat(output) -@amp.autocast(enabled=False) +@amp.autocast("cuda", enabled=False) def rope_apply_multires(x, x_lens, x_shapes, freqs, pad=True): """ x: [B*L, N, C]. diff --git a/scepter/modules/model/backbone/unet/unet_module.py b/scepter/modules/model/backbone/unet/unet_module.py index 67ccb27..0b8cb7b 100644 --- a/scepter/modules/model/backbone/unet/unet_module.py +++ b/scepter/modules/model/backbone/unet/unet_module.py @@ -459,7 +459,7 @@ class DiffusionUNet(BaseModel): from safetensors.torch import load_file as load_safetensors sd = load_safetensors(path) else: - sd = torch.load(path, map_location='cpu') + sd = torch.load(path, map_location='cpu', weights_only=True) new_sd = OrderedDict() for k, v in sd.items(): @@ -1231,7 +1231,7 @@ class LargenUNetXL(DiffusionUNetXL): from safetensors.torch import load_file as load_safetensors sd = load_safetensors(path) else: - sd = torch.load(path, map_location='cpu') + sd = torch.load(path, map_location='cpu', weights_only=True) new_sd = OrderedDict() for k, v in sd.items(): diff --git a/scepter/modules/model/diffusion/__init__.py b/scepter/modules/model/diffusion/__init__.py index 3671486..14414c5 100644 --- a/scepter/modules/model/diffusion/__init__.py +++ b/scepter/modules/model/diffusion/__init__.py @@ -1,7 +1,27 @@ # -*- coding: utf-8 -*- # Copyright (c) Alibaba, Inc. and its affiliates. +from typing import TYPE_CHECKING +from scepter.modules.utils.import_utils import LazyImportModule -from .diffusions import BaseDiffusion, DiffusionFluxRF -from .samplers import BaseDiffusionSampler, DDIMSampler, FlowEluerSampler -from .schedules import (BaseNoiseScheduler, FlowMatchShiftScheduler, - ScaledLinearScheduler) + +if TYPE_CHECKING: + from .diffusions import BaseDiffusion, DiffusionFluxRF + from .samplers import BaseDiffusionSampler, DDIMSampler, FlowEluerSampler + from .schedules import (BaseNoiseScheduler, FlowMatchShiftScheduler, + ScaledLinearScheduler) +else: + _import_structure = { + 'diffusions': ['BaseDiffusion', 'DiffusionFluxRF'], + 'samplers': ['BaseDiffusionSampler', 'DDIMSampler', 'FlowEluerSampler'], + 'schedules': ['BaseNoiseScheduler', 'FlowMatchShiftScheduler', + 'ScaledLinearScheduler'] + } + + import sys + sys.modules[__name__] = LazyImportModule( + __name__, + globals()['__file__'], + _import_structure, + module_spec=__spec__, + extra_objects={}, + ) diff --git a/scepter/modules/model/diffusion/diffusions.py b/scepter/modules/model/diffusion/diffusions.py index 7c7e15f..1f1fbe0 100644 --- a/scepter/modules/model/diffusion/diffusions.py +++ b/scepter/modules/model/diffusion/diffusions.py @@ -34,6 +34,7 @@ class BaseDiffusion(object): def init_params(self): self.prediction_type = self.cfg.get('PREDICTION_TYPE', 'eps') + self.use_dynamic_cfg = self.cfg.get('USE_DYNAMIC_CFG', False) self.noise_scheduler = NOISE_SCHEDULERS.build(self.cfg.NOISE_SCHEDULER, logger=self.logger) self.sampler_scheduler = NOISE_SCHEDULERS.build(self.cfg.get( @@ -56,7 +57,6 @@ class BaseDiffusion(object): model_kwargs={}, steps=20, sampler=None, - use_dynamic_cfg=False, guide_scale=None, guide_rescale=None, show_progress=False, @@ -79,7 +79,7 @@ class BaseDiffusion(object): if guide_scale is None or guide_scale == 1.0: out = model(x=x_t, t=t, **model_kwargs) else: - if use_dynamic_cfg: + if self.use_dynamic_cfg: guidance_scale = 1 + guide_scale * ( (1 - math.cos(math.pi * ( (steps - timestamp.item()) / steps)**5.0)) / 2) @@ -158,14 +158,16 @@ class BaseDiffusion(object): def get_sampler(self, sampler): if isinstance(sampler, str): - if sampler not in DIFFUSION_SAMPLERS.class_map: + from scepter.modules.utils.import_utils import LazyImportModule + if (not LazyImportModule.get_module_type(('DIFFUSION_SAMPLERS', sampler))) and ( + sampler not in DIFFUSION_SAMPLERS.class_map): if self.logger is not None: self.logger.info( - f'{sampler} not in the defined samplers list {DIFFUSION_SAMPLERS.class_map.keys()}' + f'{sampler} not in the defined samplers list.' ) else: print( - f'{sampler} not in the defined samplers list {DIFFUSION_SAMPLERS.class_map.keys()}' + f'{sampler} not in the defined samplers list.' ) return None sampler_cfg = Config(cfg_dict={'NAME': sampler}, load=False) diff --git a/scepter/modules/model/diffusion/schedules.py b/scepter/modules/model/diffusion/schedules.py index eaef19c..e403052 100644 --- a/scepter/modules/model/diffusion/schedules.py +++ b/scepter/modules/model/diffusion/schedules.py @@ -1,6 +1,7 @@ # -*- coding: utf-8 -*- # Copyright (c) Alibaba, Inc. and its affiliates. import math +import random from dataclasses import dataclass, field from typing import Callable @@ -30,7 +31,7 @@ class ScheduleOutput(object): @NOISE_SCHEDULERS.register_class() class BaseNoiseScheduler(object): - ''' + r''' In the diffusion model, the parameters related to the noise schedule are alpha, beta, and sigma. The following are the definitions of the above three parameters, which should be the basic property for the instance of noise scheduler. @@ -483,6 +484,14 @@ class FlowMatchFluxShiftScheduler(FlowMatchUniformScheduler): 'MAX_SHIFT': { 'value': 1.15, 'description': 'The max shift factor for the timestamp.' + }, + 'PRE_T_SAMPLE': { + 'value': False, + 'description': 'Use pre-sampled timesteps or not, default is False.' + }, + 'PRE_T_SAMPLE_FOLD': { + 'value': 1, + 'description': 'The folds of pre-sampled timesteps.' } } @@ -492,6 +501,23 @@ class FlowMatchFluxShiftScheduler(FlowMatchUniformScheduler): self.sigmoid_scale = self.cfg.get('SIGMOID_SCALE', 1) self.base_shift = self.cfg.get('BASE_SHIFT', 0.5) self.max_shift = self.cfg.get('MAX_SHIFT', 1.15) + self.pre_t_sample = self.cfg.get('PRE_T_SAMPLE', False) + self.pre_t_sample_fold = self.cfg.get('PRE_T_SAMPLE_FOLD', 1) + if self.pre_t_sample: + t = torch.sigmoid(torch.randn((self.num_timesteps * self.pre_t_sample_fold,))) + # Scale and reverse the values to go from 1000 to 0 + timesteps = ((1 - t) * 1000) + # Sort the timesteps in descending order + self.pre_sample_timesteps, _ = torch.sort(timesteps, descending=True) + else: + self.pre_sample_timesteps = None + + @property + def pre_timesteps(self): + fold_id = random.randint(0, self.pre_t_sample_fold - 1) + # print("fold_id", fold_id) + return self.pre_sample_timesteps[fold_id::self.pre_t_sample_fold] + def time_shift(self, mu: float, sigma_scale: float, t: Tensor): return math.exp(mu) / (math.exp(mu) + (1 / t - 1)**sigma_scale) @@ -516,11 +542,22 @@ class FlowMatchFluxShiftScheduler(FlowMatchUniformScheduler): n, _, h, w = x_0.shape seq_len = (h // 2 * w // 2) if t is None: - logits_norm = torch.randn(x_0.shape[0], device=x_0.device) - logits_norm = logits_norm * self.sigmoid_scale # larger scale for more uniform sampling - t = logits_norm.sigmoid() * self.num_timesteps + if self.pre_t_sample: + timestep_indices = torch.randint( + 1, + self.num_timesteps - 1, + (x_0.shape[0],) + ) + timestep_indices = timestep_indices.long() + t = [self.pre_timesteps[x.item()].to(x_0.device) for x in timestep_indices] + t = torch.stack(t, dim=0) + else: + logits_norm = torch.randn(x_0.shape[0], device=x_0.device) + logits_norm = logits_norm * self.sigmoid_scale # larger scale for more uniform sampling + t = logits_norm.sigmoid() * self.num_timesteps sigma = self.t_to_sigma(t, seq_len=seq_len) shape = (x_0.size(0), ) + (1, ) * (x_0.ndim - 1) + # print(sigma) x_t = (1 - sigma.view(shape)) * x_0 + sigma.view(shape) * noise return ScheduleOutput(x_0=x_0, x_t=x_t, diff --git a/scepter/modules/model/embedder/__init__.py b/scepter/modules/model/embedder/__init__.py index 88524d8..9a276e2 100644 --- a/scepter/modules/model/embedder/__init__.py +++ b/scepter/modules/model/embedder/__init__.py @@ -1,8 +1,30 @@ # -*- coding: utf-8 -*- # Copyright (c) Alibaba, Inc. and its affiliates. +from typing import TYPE_CHECKING +from scepter.modules.utils.import_utils import LazyImportModule -from scepter.modules.model.embedder.embedder import ( - ConcatTimestepEmbedderND, FrozenCLIPEmbedder, FrozenCLIPEmbedder2, - FrozenOpenCLIPEmbedder, FrozenOpenCLIPEmbedder2, GeneralConditioner, - IPAdapterPlusEmbedder, RefCrossEmbedder, SD3TextEmbedder, T5EmbedderHF) -from scepter.modules.model.embedder.flux_embedder import HFEmbedder + +if TYPE_CHECKING: + from scepter.modules.model.embedder.embedder import ( + ConcatTimestepEmbedderND, FrozenCLIPEmbedder, FrozenCLIPEmbedder2, + FrozenOpenCLIPEmbedder, FrozenOpenCLIPEmbedder2, GeneralConditioner, + IPAdapterPlusEmbedder, RefCrossEmbedder, SD3TextEmbedder, T5EmbedderHF) + from scepter.modules.model.embedder.flux_embedder import HFEmbedder +else: + _import_structure = { + 'embedder': ['ConcatTimestepEmbedderND', 'FrozenCLIPEmbedder', + 'FrozenCLIPEmbedder2', 'FrozenOpenCLIPEmbedder', + 'FrozenOpenCLIPEmbedder2', 'GeneralConditioner', + 'IPAdapterPlusEmbedder', 'RefCrossEmbedder', + 'SD3TextEmbedder', 'T5EmbedderHF'], + 'flux_embedder': ['HFEmbedder'] + } + + import sys + sys.modules[__name__] = LazyImportModule( + __name__, + globals()['__file__'], + _import_structure, + module_spec=__spec__, + extra_objects={}, + ) diff --git a/scepter/modules/model/embedder/embedder.py b/scepter/modules/model/embedder/embedder.py index 4d9f563..9b54672 100644 --- a/scepter/modules/model/embedder/embedder.py +++ b/scepter/modules/model/embedder/embedder.py @@ -36,7 +36,8 @@ except Exception as e: def autocast(f, enabled=True): def do_autocast(*args, **kwargs): - with torch.cuda.amp.autocast( + with torch.amp.autocast( + "cuda", enabled=enabled, dtype=torch.get_autocast_gpu_dtype(), cache_enabled=torch.is_autocast_cache_enabled(), @@ -239,7 +240,7 @@ class FrozenOpenCLIPEmbedder(BaseEmbedder): if cfg.PRETRAINED_MODEL is not None: with FS.get_from(cfg.PRETRAINED_MODEL, wait_finish=True) as local_path: - model.load_state_dict(torch.load(local_path), strict=False) + model.load_state_dict(torch.load(local_path, weights_only=True), strict=False) self.model = model self.use_grad = cfg.get('USE_GRAD', False) @@ -538,7 +539,7 @@ class IPAdapterPlusEmbedder(BaseEmbedder): ) with FS.get_from(cfg.PRETRAINED_MODEL, wait_finish=True) as local_path: - ckpt = torch.load(local_path, map_location='cpu') + ckpt = torch.load(local_path, map_location='cpu', weights_only=True) self.image_proj_model.load_state_dict(ckpt['image_proj'], strict=True) @@ -645,7 +646,7 @@ class GeneralConditioner(BaseEmbedder): from safetensors.torch import load_file as load_safetensors sd = load_safetensors(path) else: - sd = torch.load(path, map_location='cpu') + sd = torch.load(path, map_location='cpu', weights_only=True) new_sd = OrderedDict() for k, v in sd.items(): ignored = False diff --git a/scepter/modules/model/embedder/flux_embedder.py b/scepter/modules/model/embedder/flux_embedder.py index 7ab56c7..6dc7b9d 100644 --- a/scepter/modules/model/embedder/flux_embedder.py +++ b/scepter/modules/model/embedder/flux_embedder.py @@ -1,5 +1,7 @@ # -*- coding: utf-8 -*- # Copyright (c) Alibaba, Inc. and its affiliates. +# This file contains code that is adapted from +# https://github.com/black-forest-labs/flux.git import torch import transformers from scepter.modules.model.embedder.base_embedder import BaseEmbedder @@ -51,59 +53,53 @@ class HFEmbedder(BaseEmbedder): def __init__(self, cfg, logger=None): super().__init__(cfg, logger=logger) hf_model_cls = cfg.get('HF_MODEL_CLS', None) - model_path = cfg.get('MODEL_PATH', None) + model_path = cfg.get("MODEL_PATH", None) hf_tokenizer_cls = cfg.get('HF_TOKENIZER_CLS', None) tokenizer_path = cfg.get('TOKENIZER_PATH', None) self.max_length = cfg.get('MAX_LENGTH', 77) - self.output_key = cfg.get('OUTPUT_KEY', 'last_hidden_state') - self.d_type = cfg.get('D_TYPE', 'float') - self.clean = cfg.get('CLEAN', 'whitespace') - self.batch_infer = cfg.get('BATCH_INFER', False) + self.output_key = cfg.get("OUTPUT_KEY", "last_hidden_state") + self.d_type = cfg.get("D_TYPE", "float") + self.clean = cfg.get("CLEAN", "whitespace") + self.batch_infer = cfg.get("BATCH_INFER", False) + self.added_identifier = cfg.get('ADDED_IDENTIFIER', None) torch_dtype = getattr(torch, self.d_type) assert hf_model_cls is not None and hf_tokenizer_cls is not None assert model_path is not None and tokenizer_path is not None + with FS.get_dir_to_local_dir(tokenizer_path, wait_finish=True) as local_path: + self.tokenizer = getattr(transformers, hf_tokenizer_cls).from_pretrained(local_path, + max_length = self.max_length, + torch_dtype = torch_dtype, + additional_special_tokens=self.added_identifier) - with FS.get_dir_to_local_dir(tokenizer_path, - wait_finish=True) as local_path: - self.tokenizer = getattr(transformers, - hf_tokenizer_cls).from_pretrained( - local_path, - max_length=self.max_length, - torch_dtype=torch_dtype) + with FS.get_dir_to_local_dir(model_path, wait_finish=True) as local_path: + self.hf_module = getattr(transformers, hf_model_cls).from_pretrained(local_path, torch_dtype = torch_dtype) - with FS.get_dir_to_local_dir(model_path, - wait_finish=True) as local_path: - self.hf_module = getattr(transformers, - hf_model_cls).from_pretrained( - local_path, torch_dtype=torch_dtype) self.hf_module = self.hf_module.eval().requires_grad_(False) - def forward(self, text: list[str], return_mask=False): + def forward(self, text: list[str], return_mask = False): batch_encoding = self.tokenizer( text, truncation=True, max_length=self.max_length, return_length=False, return_overflowing_tokens=False, - padding='max_length', - return_tensors='pt', + padding="max_length", + return_tensors="pt", ) outputs = self.hf_module( - input_ids=batch_encoding['input_ids'].to(self.hf_module.device), + input_ids=batch_encoding["input_ids"].to(self.hf_module.device), attention_mask=None, output_hidden_states=False, ) if return_mask: - return outputs[ - self.output_key], batch_encoding['attention_mask'].to( - self.hf_module.device) + return outputs[self.output_key], batch_encoding['attention_mask'].to(self.hf_module.device) else: return outputs[self.output_key], None - def encode(self, text, return_mask=False): + def encode(self, text, return_mask = False): if isinstance(text, str): text = [text] if self.clean: @@ -119,12 +115,36 @@ class HFEmbedder(BaseEmbedder): else: return torch.cat(cont, dim=0) else: - ret_data = self(text, return_mask=return_mask) + ret_data = self(text, return_mask = return_mask) if return_mask: return ret_data else: return ret_data[0] + def encode_list(self, text_list, return_mask=True): + cont_list = [] + mask_list = [] + for pp in text_list: + cont = self.encode(pp, return_mask=return_mask) + cont_list.append(cont[0]) if return_mask else cont_list.append(cont) + mask_list.append(cont[1]) if return_mask else mask_list.append(None) + if return_mask: + return cont_list, mask_list + else: + return cont_list + + def encode_list_of_list(self, text_list, return_mask=True): + cont_list = [] + mask_list = [] + for pp in text_list: + cont = self.encode_list(pp, return_mask=return_mask) + cont_list.append(cont[0]) if return_mask else cont_list.append(cont) + mask_list.append(cont[1]) if return_mask else mask_list.append(None) + if return_mask: + return cont_list, mask_list + else: + return cont_list + def _clean(self, text): if self.clean == 'whitespace': text = whitespace_clean(basic_clean(text)) @@ -133,7 +153,6 @@ class HFEmbedder(BaseEmbedder): elif self.clean == 'canonicalize': text = canonicalize(basic_clean(text)) return text - @staticmethod def get_config_template(): return dict_to_yaml('EMBEDDER', @@ -141,28 +160,49 @@ class HFEmbedder(BaseEmbedder): HFEmbedder.para_dict, set_name=True) - @EMBEDDERS.register_class() class T5PlusClipFluxEmbedder(BaseEmbedder): """ Uses the OpenCLIP transformer encoder for text """ - para_dict = {'T5_MODEL': {}, 'CLIP_MODEL': {}} + para_dict = { + 'T5_MODEL': {}, + 'CLIP_MODEL': {} + } def __init__(self, cfg, logger=None): super().__init__(cfg, logger=logger) self.t5_model = EMBEDDERS.build(cfg.T5_MODEL, logger=logger) self.clip_model = EMBEDDERS.build(cfg.CLIP_MODEL, logger=logger) - def encode(self, text): - t5_embeds = self.t5_model.encode(text, return_mask=False) - clip_embeds = self.clip_model.encode(text, return_mask=False) + def encode(self, text, return_mask = False): + t5_embeds = self.t5_model.encode(text, return_mask = return_mask) + clip_embeds = self.clip_model.encode(text, return_mask = return_mask) # change embedding strategy here return { 'context': t5_embeds, 'y': clip_embeds, } + def encode_list(self, text, return_mask = False): + t5_embeds = self.t5_model.encode_list(text, return_mask = return_mask) + clip_embeds = self.clip_model.encode_list(text, return_mask = return_mask) + # change embedding strategy here + return { + 'context': t5_embeds, + 'y': clip_embeds, + } + + def encode_list_of_list(self, text, return_mask = False): + t5_embeds = self.t5_model.encode_list_of_list(text, return_mask = return_mask) + clip_embeds = self.clip_model.encode_list_of_list(text, return_mask = return_mask) + # change embedding strategy here + return { + 'context': t5_embeds, + 'y': clip_embeds, + } + + @staticmethod def get_config_template(): return dict_to_yaml('EMBEDDER', diff --git a/scepter/modules/model/head/__init__.py b/scepter/modules/model/head/__init__.py index 9f70332..cc74fd7 100644 --- a/scepter/modules/model/head/__init__.py +++ b/scepter/modules/model/head/__init__.py @@ -1,5 +1,25 @@ # -*- coding: utf-8 -*- # Copyright (c) Alibaba, Inc. and its affiliates. -from scepter.modules.model.head.classifier_head import ( - ClassifierHead, CosineLinearHead, TransformerHead, TransformerHeadx2, - VideoClassifierHead, VideoClassifierHeadx2) +from typing import TYPE_CHECKING +from scepter.modules.utils.import_utils import LazyImportModule + + +if TYPE_CHECKING: + from scepter.modules.model.head.classifier_head import ( + ClassifierHead, CosineLinearHead, TransformerHead, TransformerHeadx2, + VideoClassifierHead, VideoClassifierHeadx2) +else: + _import_structure = { + 'classifier_head': ['ClassifierHead', 'CosineLinearHead', + 'TransformerHead', 'TransformerHeadx2', + 'VideoClassifierHead', 'VideoClassifierHeadx2'] + } + + import sys + sys.modules[__name__] = LazyImportModule( + __name__, + globals()['__file__'], + _import_structure, + module_spec=__spec__, + extra_objects={}, + ) diff --git a/scepter/modules/model/loss/__init__.py b/scepter/modules/model/loss/__init__.py index bdc7778..8cb92b0 100644 --- a/scepter/modules/model/loss/__init__.py +++ b/scepter/modules/model/loss/__init__.py @@ -1,4 +1,23 @@ # -*- coding: utf-8 -*- # Copyright (c) Alibaba, Inc. and its affiliates. -from scepter.modules.model.loss.base_losses import CrossEntropy -from scepter.modules.model.loss.rec_loss import MinSNRLoss, ReconstructLoss +from typing import TYPE_CHECKING +from scepter.modules.utils.import_utils import LazyImportModule + + +if TYPE_CHECKING: + from scepter.modules.model.loss.base_losses import CrossEntropy + from scepter.modules.model.loss.rec_loss import MinSNRLoss, ReconstructLoss +else: + _import_structure = { + 'base_losses': ['CrossEntropy'], + 'rec_loss': ['MinSNRLoss', 'ReconstructLoss'] + } + + import sys + sys.modules[__name__] = LazyImportModule( + __name__, + globals()['__file__'], + _import_structure, + module_spec=__spec__, + extra_objects={}, + ) diff --git a/scepter/modules/model/metric/__init__.py b/scepter/modules/model/metric/__init__.py index 155ba54..1308194 100644 --- a/scepter/modules/model/metric/__init__.py +++ b/scepter/modules/model/metric/__init__.py @@ -1,5 +1,23 @@ # -*- coding: utf-8 -*- # Copyright (c) Alibaba, Inc. and its affiliates. -from scepter.modules.model.metric.classification import (AccuracyMetric, - EnsembleAccuracyMetric - ) +from typing import TYPE_CHECKING +from scepter.modules.utils.import_utils import LazyImportModule + + +if TYPE_CHECKING: + from scepter.modules.model.metric.classification import (AccuracyMetric, + EnsembleAccuracyMetric + ) +else: + _import_structure = { + 'classification': ['AccuracyMetric', 'EnsembleAccuracyMetric'] + } + + import sys + sys.modules[__name__] = LazyImportModule( + __name__, + globals()['__file__'], + _import_structure, + module_spec=__spec__, + extra_objects={}, + ) diff --git a/scepter/modules/model/neck/__init__.py b/scepter/modules/model/neck/__init__.py index 258155d..45a6d75 100644 --- a/scepter/modules/model/neck/__init__.py +++ b/scepter/modules/model/neck/__init__.py @@ -1,5 +1,24 @@ # -*- coding: utf-8 -*- # Copyright (c) Alibaba, Inc. and its affiliates. -from scepter.modules.model.neck.global_average_pooling import \ - GlobalAveragePooling -from scepter.modules.model.neck.identity import Identity +from typing import TYPE_CHECKING +from scepter.modules.utils.import_utils import LazyImportModule + + +if TYPE_CHECKING: + from scepter.modules.model.neck.global_average_pooling import \ + GlobalAveragePooling + from scepter.modules.model.neck.identity import Identity +else: + _import_structure = { + 'global_average_pooling': ['GlobalAveragePooling'], + 'identity': ['Identity'] + } + + import sys + sys.modules[__name__] = LazyImportModule( + __name__, + globals()['__file__'], + _import_structure, + module_spec=__spec__, + extra_objects={}, + ) diff --git a/scepter/modules/model/network/__init__.py b/scepter/modules/model/network/__init__.py index 3ab7bf7..08eb717 100644 --- a/scepter/modules/model/network/__init__.py +++ b/scepter/modules/model/network/__init__.py @@ -1,9 +1,32 @@ # -*- coding: utf-8 -*- # Copyright (c) Alibaba, Inc. and its affiliates. -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_edit, ldm_pixart, - ldm_sce, ldm_sd3, ldm_xl, - ldm_flux) \ No newline at end of file +from typing import TYPE_CHECKING +from scepter.modules.utils.import_utils import LazyImportModule + + +if TYPE_CHECKING: + 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_edit, ldm_pixart, + ldm_sce, ldm_sd3, ldm_xl, + ldm_flux) +else: + _import_structure = { + 'autoencoder': ['ae_kl'], + 'classifier': ['Classifier'], + 'diffusion': ['diffusion', 'schedules', 'solvers'], + 'ldm': ['ldm', 'ldm_edit', 'ldm_pixart', + 'ldm_sce', 'ldm_sd3', 'ldm_xl', + 'ldm_flux'] + } + + import sys + sys.modules[__name__] = LazyImportModule( + __name__, + globals()['__file__'], + _import_structure, + module_spec=__spec__, + extra_objects={}, + ) diff --git a/scepter/modules/model/network/autoencoder/__init__.py b/scepter/modules/model/network/autoencoder/__init__.py index d7a10b5..afdd881 100644 --- a/scepter/modules/model/network/autoencoder/__init__.py +++ b/scepter/modules/model/network/autoencoder/__init__.py @@ -1,4 +1,23 @@ # -*- coding: utf-8 -*- # Copyright (c) Alibaba, Inc. and its affiliates. -from scepter.modules.model.network.autoencoder.ae_kl import AutoencoderKL -from scepter.modules.model.network.autoencoder.ae_kl_cogvideox import AutoencoderKLCogVideoX +from typing import TYPE_CHECKING +from scepter.modules.utils.import_utils import LazyImportModule + + +if TYPE_CHECKING: + from scepter.modules.model.network.autoencoder.ae_kl import AutoencoderKL + from scepter.modules.model.network.autoencoder.ae_kl_cogvideox import AutoencoderKLCogVideoX +else: + _import_structure = { + 'ae_kl': ['AutoencoderKL'], + 'ae_kl_cogvideox': ['AutoencoderKLCogVideoX'] + } + + import sys + sys.modules[__name__] = LazyImportModule( + __name__, + globals()['__file__'], + _import_structure, + module_spec=__spec__, + extra_objects={}, + ) diff --git a/scepter/modules/model/network/autoencoder/ae_kl.py b/scepter/modules/model/network/autoencoder/ae_kl.py index 3d0fc95..58e13ca 100644 --- a/scepter/modules/model/network/autoencoder/ae_kl.py +++ b/scepter/modules/model/network/autoencoder/ae_kl.py @@ -129,7 +129,7 @@ class AutoencoderKL(TrainModule): for k in f.keys(): sd[k] = f.get_tensor(k) else: - sd = torch.load(path, map_location='cpu') + sd = torch.load(path, map_location='cpu', weights_only=True) if path.find('.pt') > -1 and 'state_dict' in sd: sd = sd['state_dict'] elif path.find('.ckpt') > -1 and 'state_dict' in sd: @@ -373,7 +373,7 @@ class AutoencoderKLFlux(TrainModule): for k in f.keys(): sd[k] = f.get_tensor(k) else: - sd = torch.load(path, map_location="cpu") + sd = torch.load(path, map_location="cpu", weights_only=True) if path.find('.pt') > -1 and 'state_dict' in sd: sd = sd['state_dict'] elif path.find('.ckpt') > -1 and 'state_dict' in sd: diff --git a/scepter/modules/model/network/autoencoder/ae_kl_cogvideox.py b/scepter/modules/model/network/autoencoder/ae_kl_cogvideox.py index f947351..22006d5 100644 --- a/scepter/modules/model/network/autoencoder/ae_kl_cogvideox.py +++ b/scepter/modules/model/network/autoencoder/ae_kl_cogvideox.py @@ -1591,7 +1591,7 @@ class AutoencoderKLCogVideoX(TrainModule): from safetensors.torch import load_file as load_safetensors ckpt = load_safetensors(local_model) else: - ckpt = torch.load(local_model, map_location='cpu') + ckpt = torch.load(local_model, map_location='cpu', weights_only=True) missing, unexpected = self.load_state_dict(ckpt, strict=False) if we.rank == 0: self.logger.info( diff --git a/scepter/modules/model/network/diffusion/__init__.py b/scepter/modules/model/network/diffusion/__init__.py index f00e199..ee6adb7 100644 --- a/scepter/modules/model/network/diffusion/__init__.py +++ b/scepter/modules/model/network/diffusion/__init__.py @@ -1,4 +1,22 @@ # -*- coding: utf-8 -*- # Copyright (c) Alibaba, Inc. and its affiliates. -from scepter.modules.model.network.diffusion import (diffusion, schedules, - solvers) +from typing import TYPE_CHECKING +from scepter.modules.utils.import_utils import LazyImportModule + + +if TYPE_CHECKING: + from scepter.modules.model.network.diffusion import (diffusion, schedules, + solvers) +else: + _import_structure = { + 'diffusion': ['diffusion', 'schedules', 'solvers'] + } + + import sys + sys.modules[__name__] = LazyImportModule( + __name__, + globals()['__file__'], + _import_structure, + module_spec=__spec__, + extra_objects={}, + ) diff --git a/scepter/modules/model/network/diffusion/diffusion.py b/scepter/modules/model/network/diffusion/diffusion.py index 1ea84f0..e78bb2f 100644 --- a/scepter/modules/model/network/diffusion/diffusion.py +++ b/scepter/modules/model/network/diffusion/diffusion.py @@ -239,7 +239,7 @@ class GaussianDiffusion(object): percentile=None, cat_uc=False, **kwargs): - """ + r""" Apply one step of denoising from the posterior distribution q(x_s | x_t, x0). Since x0 is not available, estimate the denoising results using the learned distribution p(x_s | x_t, \hat{x}_0 == f(x_t)). # noqa diff --git a/scepter/modules/model/network/ldm/__init__.py b/scepter/modules/model/network/ldm/__init__.py index 9ebba17..946fde9 100644 --- a/scepter/modules/model/network/ldm/__init__.py +++ b/scepter/modules/model/network/ldm/__init__.py @@ -1,15 +1,44 @@ # -*- coding: utf-8 -*- # Copyright (c) Alibaba, Inc. and its affiliates. -from scepter.modules.model.network.ldm.ldm import LatentDiffusion -from scepter.modules.model.network.ldm.ldm_ace import (LatentDiffusionACE, - LatentDiffusionACERefiner) -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 -from scepter.modules.model.network.ldm.ldm_cogvideox import LatentDiffusionCogVideoX -from scepter.modules.model.network.ldm.ldm_flux import (LatentDiffusionFlux, - LatentDiffusionFluxMR) +from typing import TYPE_CHECKING +from scepter.modules.utils.import_utils import LazyImportModule + + +if TYPE_CHECKING: + from scepter.modules.model.network.ldm.ldm import LatentDiffusion + from scepter.modules.model.network.ldm.ldm_ace import (LatentDiffusionACE, + LatentDiffusionACERefiner) + 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 + from scepter.modules.model.network.ldm.ldm_cogvideox import LatentDiffusionCogVideoX + from scepter.modules.model.network.ldm.ldm_flux import (LatentDiffusionFlux, + LatentDiffusionFluxMR) + from scepter.modules.model.network.ldm.ldm_ace_plus import LatentDiffusionACEPlus +else: + _import_structure = { + 'ldm': ['LatentDiffusion'], + 'ldm_ace': ['LatentDiffusionACE', 'LatentDiffusionACERefiner'], + 'ldm_edit': ['LatentDiffusionEdit'], + 'ldm_pixart': ['LatentDiffusionPixart'], + 'ldm_sce': ['LatentDiffusionSCEControl', 'LatentDiffusionSCETuning', + 'LatentDiffusionXLSCEControl', 'LatentDiffusionXLSCETuning'], + 'ldm_sd3': ['LatentDiffusionSD3'], + 'ldm_xl': ['LatentDiffusionXL'], + 'ldm_cogvideox': ['LatentDiffusionCogVideoX'], + 'ldm_flux': ['LatentDiffusionFlux', 'LatentDiffusionFluxMR'], + 'ldm_ace_plus': ['LatentDiffusionACEPlus'], + } + + import sys + sys.modules[__name__] = LazyImportModule( + __name__, + globals()['__file__'], + _import_structure, + module_spec=__spec__, + extra_objects={}, + ) diff --git a/scepter/modules/model/network/ldm/ldm.py b/scepter/modules/model/network/ldm/ldm.py index 5c98b2b..e199125 100644 --- a/scepter/modules/model/network/ldm/ldm.py +++ b/scepter/modules/model/network/ldm/ldm.py @@ -200,7 +200,7 @@ class LatentDiffusion(TrainModule): from safetensors.torch import load_file as load_safetensors sd = load_safetensors(path) else: - sd = torch.load(path, map_location='cpu') + sd = torch.load(path, map_location='cpu',weights_only=True) new_sd = OrderedDict() for k, v in sd.items(): ignored = False diff --git a/scepter/modules/model/network/ldm/ldm_ace.py b/scepter/modules/model/network/ldm/ldm_ace.py index 215e156..fdd6192 100644 --- a/scepter/modules/model/network/ldm/ldm_ace.py +++ b/scepter/modules/model/network/ldm/ldm_ace.py @@ -95,8 +95,8 @@ class LatentDiffusionACE(LatentDiffusion): return batch_data_list def forward_train(self, - edit_image=[], - edit_image_mask=[], + src_image_list=[], + src_mask_list=[], image=None, image_mask=None, noise=None, @@ -114,8 +114,8 @@ class LatentDiffusionACE(LatentDiffusion): Returns: ''' assert check_list_of_list(prompt) and check_list_of_list( - edit_image) and check_list_of_list(edit_image_mask) - assert len(edit_image) == len(edit_image_mask) == len(prompt) + src_image_list) and check_list_of_list(src_mask_list) + assert len(src_image_list) == len(src_mask_list) == len(prompt) assert self.cond_stage_model is not None gc_seg = kwargs.pop('gc_seg', []) gc_seg = int(gc_seg[0]) if len(gc_seg) > 0 else 0 @@ -143,13 +143,13 @@ class LatentDiffusionACE(LatentDiffusion): 'encode_list_of_list')(prompt_, return_mask=True) except Exception as e: print(e, prompt_) - cont, cont_mask = self.cond_stage_embeddings(prompt, edit_image, cont, + cont, cont_mask = self.cond_stage_embeddings(prompt, src_image_list, cont, cont_mask) context['crossattn'] = cont # process edit image & edit image mask - edit_image = [to_device(i, strict=False) for i in edit_image] - edit_image_mask = [to_device(i, strict=False) for i in edit_image_mask] + edit_image = [to_device(i, strict=False) for i in src_image_list] + edit_image_mask = [to_device(i, strict=False) for i in src_mask_list] e_img, e_mask = [], [] for u, m in zip(edit_image, edit_image_mask): if m is None: @@ -185,8 +185,8 @@ class LatentDiffusionACE(LatentDiffusion): @torch.no_grad() def forward_test(self, - edit_image=[], - edit_image_mask=[], + src_image_list=[], + src_mask_list=[], image=None, image_mask=None, prompt=[], @@ -200,8 +200,8 @@ class LatentDiffusionACE(LatentDiffusion): **kwargs): assert check_list_of_list(prompt) and check_list_of_list( - edit_image) and check_list_of_list(edit_image_mask) - assert len(edit_image) == len(edit_image_mask) == len(prompt) + src_image_list) and check_list_of_list(src_mask_list) + assert len(src_image_list) == len(src_mask_list) == len(prompt) assert self.cond_stage_model is not None # gc_seg is unused kwargs.pop('gc_seg', -1) @@ -209,7 +209,7 @@ class LatentDiffusionACE(LatentDiffusion): context, null_context = {}, {} prompt, n_prompt, image, image_mask, edit_image, edit_image_mask = self.limit_batch_data( - [prompt, n_prompt, image, image_mask, edit_image, edit_image_mask], + [prompt, n_prompt, image, image_mask, src_image_list, src_mask_list], log_num) g = torch.Generator(device=we.device_id) seed = seed if seed >= 0 else random.randint(0, 2**32 - 1) @@ -368,8 +368,8 @@ class LatentDiffusionACERefiner(LatentDiffusionACE): self.enhence_sampler_cfg = None def forward_sample(self, - edit_image=[], - edit_mask=[], + src_image_list=[], + src_mask_list=[], noise=None, cond_mask=[], x_shapes=[], @@ -414,15 +414,15 @@ class LatentDiffusionACERefiner(LatentDiffusionACE): # with torch.autocast(device_type="cuda", enabled=True, dtype=torch.bfloat16): cont, cont_mask = getattr(self.cond_stage_model, 'encode_list')(prompt, return_mask=True) - cont, cont_mask = self.cond_stage_embeddings(prompt, edit_image, cont, cont_mask) + cont, cont_mask = self.cond_stage_embeddings(prompt, src_image_list, cont, cont_mask) null_cont, null_cont_mask = getattr(self.cond_stage_model, 'encode_list')(n_prompt, return_mask=True) - null_cont, null_cont_mask = self.cond_stage_embeddings(prompt, edit_image, null_cont, null_cont_mask) + null_cont, null_cont_mask = self.cond_stage_embeddings(prompt, src_image_list, null_cont, null_cont_mask) context['crossattn'] = cont null_context['crossattn'] = null_cont - null_context['edit'] = context['edit'] = edit_image - null_context['edit_mask'] = context['edit_mask'] = edit_mask + null_context['edit'] = context['edit'] = src_image_list + null_context['edit_mask'] = context['edit_mask'] = src_mask_list # process sample model = self.model_ema if self.use_ema and self.eval_ema else self.model @@ -478,8 +478,8 @@ class LatentDiffusionACERefiner(LatentDiffusionACE): @torch.no_grad() def forward_test(self, - edit_image=[], - edit_image_mask=[], + src_image_list=[], + src_mask_list=[], image=None, image_mask=None, prompt=[], @@ -493,13 +493,13 @@ class LatentDiffusionACERefiner(LatentDiffusionACE): enhance_scale=0.99, log_num=-1, **kwargs): - assert check_list_of_list(prompt) and check_list_of_list(edit_image) and check_list_of_list(edit_image_mask) - assert len(edit_image) == len(edit_image_mask) == len(prompt) + assert check_list_of_list(prompt) and check_list_of_list(src_image_list) and check_list_of_list(src_mask_list) + assert len(src_image_list) == len(src_mask_list) == len(prompt) assert self.cond_stage_model is not None # gc_seg is unused kwargs.pop("gc_seg", -1) prompt, n_prompt, image, image_mask, edit_image, edit_image_mask = self.limit_batch_data( - [prompt, n_prompt, image, image_mask, edit_image, edit_image_mask], log_num) + [prompt, n_prompt, image, image_mask, src_image_list, src_mask_list], log_num) prompt = [[pp] if isinstance(pp, str) else pp for pp in prompt] diff --git a/scepter/modules/model/network/ldm/ldm_ace_plus.py b/scepter/modules/model/network/ldm/ldm_ace_plus.py new file mode 100644 index 0000000..fea6fc1 --- /dev/null +++ b/scepter/modules/model/network/ldm/ldm_ace_plus.py @@ -0,0 +1,371 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +import copy +import random +from contextlib import nullcontext + +import torch +import torch.nn.functional as F +from torch.distributed.fsdp import FullyShardedDataParallel + +from einops import rearrange +from scepter.modules.model.network.ldm import LatentDiffusionFluxMR +from scepter.modules.model.registry import MODELS +from scepter.modules.model.utils.basic_utils import ( + check_list_of_list, limit_batch_data, pack_imagelist_into_tensor, + to_device, unpack_tensor_into_imagelist) +from scepter.modules.utils.config import dict_to_yaml +from scepter.modules.utils.distribute import we + + +@MODELS.register_class() +class LatentDiffusionACEPlus(LatentDiffusionFluxMR): + para_dict = {} + para_dict.update(LatentDiffusionFluxMR.para_dict) + + def resize_func(self, x, size): + if x is None: + return x + return F.interpolate(x.unsqueeze(0), size=size, mode='nearest-exact') + + def parse_ref_and_edit( + self, + src_image, + src_image_mask, + text_embedding, + # text_mask, + edit_id): + edit_image = [] + edit_mask = [] + ref_image = [] + ref_mask = [] + ref_context = [] + ref_y = [] + ref_id = [] + txt = [] + txt_y = [] + for sample_id, ( + one_src, + one_src_mask, + one_text_embedding, + one_text_y, + # one_text_mask, + one_edit_id) in enumerate( + zip( + src_image, + src_image_mask, + text_embedding['context'], + text_embedding['y'], + # text_mask, + edit_id)): + ref_id.append([i for i in range(len(one_src))]) + if hasattr(self, + 'ref_cond_stage_model') and self.ref_cond_stage_model: + ref_image.append( + self.ref_cond_stage_model.encode_list([ + ((i + 1.0) / 2.0 * 255).type(torch.uint8) + for i in one_src + ])) + else: + ref_image.append(one_src) + ref_mask.append(one_src_mask) + # process edit image & edit image mask + current_edit_image = to_device([one_src[i] for i in one_edit_id], + strict=False) + current_edit_image = [ + v.squeeze(0) + for v in self.encode_first_stage(current_edit_image) + ] + current_edit_image_mask = to_device( + [one_src_mask[i] for i in one_edit_id], strict=False) + current_edit_image_mask = [ + self.reshape_func(m).squeeze(0) + for m in current_edit_image_mask + ] + + edit_image.append(current_edit_image) + edit_mask.append(current_edit_image_mask) + ref_context.append(one_text_embedding[:len(ref_id[-1])]) + ref_y.append(one_text_y[:len(ref_id[-1])]) + if not sum(len(src_) for src_ in src_image) > 0: + ref_image = None + ref_context = None + ref_y = None + for sample_id, (one_text_embedding, one_text_y) in enumerate( + zip(text_embedding['context'], text_embedding['y'])): + txt.append(one_text_embedding[-1].squeeze(0)) + txt_y.append(one_text_y[-1]) + return { + 'edit': edit_image, + 'edit_mask': edit_mask, + 'edit_id': edit_id, + 'ref_context': ref_context, + 'ref_y': ref_y, + 'context': txt, + 'y': txt_y, + 'ref_x': ref_image, + 'ref_mask': ref_mask, + 'ref_id': ref_id + } + + def reshape_func(self, mask): + mask = mask.to(torch.bfloat16) + mask = mask.view((-1, mask.shape[-2], mask.shape[-1])) + mask = rearrange( + mask, + 'c (h ph) (w pw) -> c (ph pw) h w', + ph=8, + pw=8, + ) + return mask + + def forward_train(self, + src_image_list=[], + src_mask_list=[], + edit_id=[], + image=None, + image_mask=None, + noise=None, + prompt=[], + **kwargs): + ''' + Args: + src_image: list of list of src_image + src_image_mask: list of list of src_image_mask + image: target image + image_mask: target image mask + noise: default is None, generate automaticly + ref_prompt: list of list of text + prompt: list of text + **kwargs: + Returns: + ''' + assert check_list_of_list(src_image_list) and check_list_of_list( + src_mask_list) + assert self.cond_stage_model is not None + + gc_seg = kwargs.pop('gc_seg', []) + gc_seg = int(gc_seg[0]) if len(gc_seg) > 0 else 0 + align = kwargs.pop('align', []) + prompt_ = [[pp] if isinstance(pp, str) else pp for pp in prompt] + if len(align) < 1: + align = [0] * len(prompt_) + context = getattr(self.cond_stage_model, + 'encode_list_of_list')(prompt_) + guide_scale = self.guide_scale + if guide_scale is not None: + guide_scale = torch.full((len(prompt_), ), + guide_scale, + device=we.device_id) + else: + guide_scale = None + # image and image_mask + # print("is list of list", check_list_of_list(image)) + if check_list_of_list(image): + image = [to_device(ix) for ix in image] + x_start = [self.encode_first_stage(ix, **kwargs) for ix in image] + noise = [[torch.randn_like(ii) for ii in ix] for ix in x_start] + x_start = [torch.cat(ix, dim=-1) for ix in x_start] + noise = [torch.cat(ix, dim=-1) for ix in noise] + + noise, _ = pack_imagelist_into_tensor(noise) + + image_mask = [to_device(im, strict=False) for im in image_mask] + x_mask = [[self.reshape_func(i).squeeze(0) + for i in im] if im is not None else [None] * len(ix) + for ix, im in zip(image, image_mask)] + x_mask = [torch.cat(im, dim=-1) for im in x_mask] + else: + image = to_device(image) + x_start = self.encode_first_stage(image, **kwargs) + image_mask = to_device(image_mask, strict=False) + x_mask = [self.reshape_func(i).squeeze(0) for i in image_mask + ] if image_mask is not None else [None] * len(image) + loss_mask, _ = pack_imagelist_into_tensor( + tuple( + torch.ones_like(ix, dtype=torch.bool, device=ix.device) + for ix in x_start)) + x_start, x_shapes = pack_imagelist_into_tensor(x_start) + context['x_shapes'] = x_shapes + context['align'] = align + # process image mask + + context['x_mask'] = x_mask + ref_edit_context = self.parse_ref_and_edit(src_image_list, + src_mask_list, context, + edit_id) + context.update(ref_edit_context) + + teacher_context = copy.deepcopy(context) + teacher_context['context'] = torch.cat(teacher_context['context'], + dim=0) + teacher_context['y'] = torch.cat(teacher_context['y'], dim=0) + loss = self.diffusion.loss(x_0=x_start, + model=self.model, + model_kwargs={ + 'cond': context, + 'gc_seg': gc_seg, + 'guidance': guide_scale + }, + noise=noise, + reduction='none', + **kwargs) + loss = loss[loss_mask].mean() + ret = {'loss': loss, 'probe_data': {'prompt': prompt}} + return ret + + @torch.no_grad() + def forward_test(self, + src_image_list=[], + src_mask_list=[], + edit_id=[], + image=None, + image_mask=None, + prompt=[], + sampler='flow_euler', + sample_steps=20, + seed=2023, + guide_scale=3.5, + guide_rescale=0.0, + show_process=False, + log_num=-1, + **kwargs): + outputs = self.forward_editing(src_image_list=src_image_list, + src_mask_list=src_mask_list, + edit_id=edit_id, + image=image, + image_mask=image_mask, + prompt=prompt, + sampler=sampler, + sample_steps=sample_steps, + seed=seed, + guide_scale=guide_scale, + guide_rescale=guide_rescale, + show_process=show_process, + log_num=log_num, + **kwargs) + return outputs + + @torch.no_grad() + def forward_editing(self, + src_image_list=[], + src_mask_list=[], + edit_id=[], + image=None, + image_mask=None, + prompt=[], + sampler='flow_euler', + sample_steps=20, + seed=2023, + guide_scale=3.5, + log_num=-1, + **kwargs): + # gc_seg is unused + prompt, image, image_mask, src_image, src_image_mask, edit_id = limit_batch_data( + [ + prompt, image, image_mask, src_image_list, src_mask_list, + edit_id + ], log_num) + assert check_list_of_list(src_image) and check_list_of_list( + src_image_mask) + assert self.cond_stage_model is not None + align = kwargs.pop('align', []) + prompt_ = [[pp] if isinstance(pp, str) else pp for pp in prompt] + if len(align) < 1: + align = [0] * len(prompt_) + context = getattr(self.cond_stage_model, + 'encode_list_of_list')(prompt_) + guide_scale = guide_scale or self.guide_scale + if guide_scale is not None: + guide_scale = torch.full((len(prompt), ), + guide_scale, + device=we.device_id) + else: + guide_scale = None + # image and image_mask + seed = seed if seed >= 0 else random.randint(0, 2**32 - 1) + if image is not None: + if check_list_of_list(image): + image = [torch.cat(ix, dim=-1) for ix in image] + image_mask = [torch.cat(im, dim=-1) for im in image_mask] + noise = [ + self.noise_sample(1, ix.shape[1], ix.shape[2], seed) + for ix in image + ] + else: + height, width = kwargs.pop('height'), kwargs.pop('width') + noise = [self.noise_sample(1, height, width, seed) for _ in prompt] + noise, x_shapes = pack_imagelist_into_tensor(noise) + context['x_shapes'] = x_shapes + context['align'] = align + # process image mask + image_mask = to_device(image_mask, strict=False) + x_mask = [self.reshape_func(i).squeeze(0) for i in image_mask] + context['x_mask'] = x_mask + ref_edit_context = self.parse_ref_and_edit(src_image, src_image_mask, + context, edit_id) + context.update(ref_edit_context) + # UNet use input n_prompt + # model = self.model_ema if self.use_ema and self.eval_ema else self.model + # import pdb;pdb.set_trace() + model = self.model + embedding_context = model.no_sync if isinstance(model, FullyShardedDataParallel) \ + else nullcontext + with embedding_context(): + samples = self.diffusion.sample(noise=noise, + sampler=sampler, + model=self.model, + model_kwargs={ + 'cond': context, + 'guidance': guide_scale, + 'gc_seg': -1 + }, + steps=sample_steps, + show_progress=True, + guide_scale=guide_scale, + return_intermediate=None, + **kwargs).float() + samples = unpack_tensor_into_imagelist(samples, x_shapes) + with torch.autocast(device_type='cuda', dtype=torch.bfloat16): + x_samples = self.decode_first_stage(samples) + outputs = list() + for i in range(len(prompt)): + rec_img = torch.clamp((x_samples[i].float() + 1.0) / 2.0, + min=0.0, + max=1.0) + rec_img = rec_img.squeeze(0) + edit_imgs, edit_img_masks = [], [] + if src_image is not None and src_image[i] is not None: + if src_image_mask[i] is None: + src_image_mask[i] = [None] * len(src_image[i]) + for edit_img, edit_mask in zip(src_image[i], + src_image_mask[i]): + edit_img = torch.clamp((edit_img.float() + 1.0) / 2.0, + min=0.0, + max=1.0) + edit_imgs.append(edit_img.squeeze(0)) + if edit_mask is None: + edit_mask = torch.ones_like(edit_img[[0], :, :]) + edit_img_masks.append(edit_mask) + one_tup = { + 'reconstruct_image': rec_img, + 'instruction': prompt[i], + 'edit_image': edit_imgs if len(edit_imgs) > 0 else None, + 'edit_mask': edit_img_masks if len(edit_imgs) > 0 else None + } + if image is not None: + if image_mask is None: + image_mask = [None] * len(image) + ori_img = torch.clamp((image[i] + 1.0) / 2.0, min=0.0, max=1.0) + one_tup['target_image'] = ori_img.squeeze(0) + one_tup['target_mask'] = image_mask[i] if image_mask[ + i] is not None else torch.ones_like(ori_img[[0], :, :]) + outputs.append(one_tup) + return outputs + + @staticmethod + def get_config_template(): + return dict_to_yaml('MODEL', + __class__.__name__, + LatentDiffusionACEPlus.para_dict, + set_name=True) diff --git a/scepter/modules/model/network/ldm/ldm_cogvideox.py b/scepter/modules/model/network/ldm/ldm_cogvideox.py index 256e24c..8335602 100644 --- a/scepter/modules/model/network/ldm/ldm_cogvideox.py +++ b/scepter/modules/model/network/ldm/ldm_cogvideox.py @@ -26,19 +26,27 @@ class LatentDiffusionCogVideoX(LatentDiffusion): self.use_rotary_positional_embeddings = self.model_config.get('USE_ROTARY_POSITIONAL_EMBEDDINGS', False) self.attention_head_dim = self.model_config.get('ATTENTION_HEAD_DIM', 64) self.patch_size = self.model_config.get('PATCH_SIZE', 2) - self.sample_height = self.first_stage_config.get('SAMPLE_HEIGHT', 480) - self.sample_width = self.first_stage_config.get('SAMPLE_WIDTH', 720) + self.patch_size_t = self.model_config.get('PATCH_SIZE_T', None) + self.ofs_embed_dim = self.model_config.get('OFS_EMBED_DIM', None) + self.sample_height = self.first_stage_config.get('SAMPLE_HEIGHT', 60) + self.sample_width = self.first_stage_config.get('SAMPLE_WIDTH', 90) self.noised_image_dropout = self.cfg.get('NOISED_IMAGE_DROPOUT', 0.05) + self.invert_scale_latents = self.cfg.get('INVERT_SCALE_LATENTS', False) def construct_network(self): super().construct_network() self.model = self.model.to(getattr(torch, self.model_config.DTYPE)) + self.first_stage_model = self.first_stage_model.to(getattr(torch, self.first_stage_config.DTYPE)) @torch.no_grad() def encode_first_stage(self, x, **kwargs): if isinstance(x, list): x = torch.stack(x, dim=0) # [B, C, F, H, W] - latents = self.scaling_factor_image * self.first_stage_model.encode(x).sample() + image_latents = self.first_stage_model.encode(x).sample() + if not self.invert_scale_latents: + latents = self.scaling_factor_image * image_latents + else: + latents = 1 / self.scaling_factor_image * image_latents return latents @torch.no_grad() @@ -78,18 +86,36 @@ class LatentDiffusionCogVideoX(LatentDiffusion): ) -> Tuple[torch.Tensor, torch.Tensor]: grid_height = height // (self.scale_factor_spatial * self.patch_size) grid_width = width // (self.scale_factor_spatial * self.patch_size) - base_size_width = self.sample_width // (self.scale_factor_spatial * self.patch_size) - base_size_height = self.sample_height // (self.scale_factor_spatial * self.patch_size) - grid_crops_coords = get_resize_crop_region_for_grid( - (grid_height, grid_width), base_size_width, base_size_height - ) - freqs_cos, freqs_sin = get_3d_rotary_pos_embed( - embed_dim=self.attention_head_dim, - crops_coords=grid_crops_coords, - grid_size=(grid_height, grid_width), - temporal_size=num_frames, - ) + p = self.patch_size + p_t = self.patch_size_t + + base_size_width = self.sample_width // p + base_size_height = self.sample_height // p + + if p_t is None: + # CogVideoX 1.0 + grid_crops_coords = get_resize_crop_region_for_grid( + (grid_height, grid_width), base_size_width, base_size_height + ) + freqs_cos, freqs_sin = get_3d_rotary_pos_embed( + embed_dim=self.attention_head_dim, + crops_coords=grid_crops_coords, + grid_size=(grid_height, grid_width), + temporal_size=num_frames, + ) + else: + # CogVideoX 1.5 + base_num_frames = (num_frames + p_t - 1) // p_t + + freqs_cos, freqs_sin = get_3d_rotary_pos_embed( + embed_dim=self.attention_head_dim, + crops_coords=None, + grid_size=(grid_height, grid_width), + temporal_size=base_num_frames, + grid_type="slice", + max_size=(base_size_height, base_size_width), + ) freqs_cos = freqs_cos.to(device=device) freqs_sin = freqs_sin.to(device=device) @@ -121,19 +147,21 @@ class LatentDiffusionCogVideoX(LatentDiffusion): else: image_latent = None - height, width = image_size + height, width = image_size[0] if isinstance(image_size, list) and all(isinstance(elem, list) for elem in image_size) else image_size image_rotary_emb = ( - self._prepare_rotary_positional_embeddings(height, width, noise.size(1), we.device_id) + self._prepare_rotary_positional_embeddings(height=height, width=width, num_frames=noise.size(1), device=we.device_id) if self.use_rotary_positional_embeddings else None ) + ofs_emb = None if self.ofs_embed_dim is None else image_latent.new_full((1,), fill_value=2.0) loss = self.diffusion.loss(x_0=x_start, t=t, model=self.model, model_kwargs={"cond": cont, 'image_latent': image_latent, - 'image_rotary_emb': image_rotary_emb}, + 'image_rotary_emb': image_rotary_emb, + 'ofs': ofs_emb}, noise=noise, **kwargs) loss = loss.mean() @@ -160,6 +188,7 @@ class LatentDiffusionCogVideoX(LatentDiffusion): image_size = [480, 720] seed = seed if seed >= 0 else random.randint(0, 2**32 - 1) generator = torch.Generator().manual_seed(seed) + # generator = torch.Generator(we.device_id).manual_seed(seed) prompt = [prompt] if isinstance(prompt, str) else prompt num_samples = len(prompt) n_prompt = default(n_prompt, [self.default_n_prompt] * len(prompt)) @@ -169,14 +198,21 @@ class LatentDiffusionCogVideoX(LatentDiffusion): cont = getattr(self.cond_stage_model, 'encode')(prompt, return_mask=False, use_mask=False) null_cont = getattr(self.cond_stage_model, 'encode')(n_prompt, return_mask=False, use_mask=False) - height, width = image_size + height, width = image_size[0] if isinstance(image_size, list) and all(isinstance(elem, list) for elem in image_size) else image_size + latent_frames = (num_frames - 1) // self.scale_factor_temporal + 1 + additional_frames = 0 + if self.patch_size_t is not None and latent_frames % self.patch_size_t != 0: + additional_frames = self.patch_size_t - latent_frames % self.patch_size_t + num_frames += additional_frames * self.scale_factor_temporal noise = self.noise_sample(num_samples, num_frames, height, width, generator) image_rotary_emb = ( self._prepare_rotary_positional_embeddings(height, width, noise.size(1), we.device_id) if self.use_rotary_positional_embeddings else None ) + image_latent, image = self.get_image_latent(image, video, noise) if image is not None else (None, None) + ofs_emb = None if self.ofs_embed_dim is None else image_latent.new_full((1,), fill_value=2.0) samples = self.diffusion.sample(noise=noise, sampler=sampler, @@ -185,19 +221,21 @@ class LatentDiffusionCogVideoX(LatentDiffusion): 'cond': cont, 'image_latent': image_latent, 'image_rotary_emb': image_rotary_emb, + 'ofs': ofs_emb }, { 'cond': null_cont, 'image_latent': image_latent, 'image_rotary_emb': image_rotary_emb, + 'ofs': ofs_emb }], steps=sample_steps, show_progress=True, - use_dynamic_cfg=True, guide_scale=guide_scale, guide_rescale=guide_rescale, return_intermediate=None, **kwargs).float() + samples = samples[:, additional_frames:] x_frames = self.decode_first_stage(samples).float() outputs = [] diff --git a/scepter/modules/model/network/ldm/ldm_flux.py b/scepter/modules/model/network/ldm/ldm_flux.py index 18455e9..aff46c8 100644 --- a/scepter/modules/model/network/ldm/ldm_flux.py +++ b/scepter/modules/model/network/ldm/ldm_flux.py @@ -14,16 +14,13 @@ from scepter.modules.model.utils.basic_utils import disabled_train, check_list_o from scepter.modules.utils.config import dict_to_yaml from scepter.modules.utils.distribute import we from scepter.modules.model.utils.basic_utils import count_params - - - @MODELS.register_class() class LatentDiffusionFlux(LatentDiffusion): para_dict = LatentDiffusion.para_dict def __init__(self, cfg, logger=None): super().__init__(cfg, logger=logger) - self.guide_scale = cfg.get('GUIDE_SCALE', 3.5) + self.guide_scale = cfg.get('GUIDE_SCALE', 1.0) def init_params(self): self.parameterization = self.cfg.get('PARAMETERIZATION', 'rf') @@ -271,7 +268,7 @@ class LatentDiffusionFluxMR(LatentDiffusionFlux): guide_scale=3.5, show_process=True, x = None, - reverse_scale = 0., + reverse_scale = -1., **kwargs ): noise, x_shapes = pack_imagelist_into_tensor(noise) @@ -377,9 +374,167 @@ class LatentDiffusionFluxMR(LatentDiffusionFlux): zu = zu[0] return zu - z = [run_one_image(u.unsqueeze(0) if u.dim == 3 else u) for u in x] + z = [run_one_image(u.unsqueeze(0) if u.dim() == 3 else u) for u in x] return z @torch.no_grad() def decode_first_stage(self, z): - return [self.first_stage_model.decode(zu) for zu in z] \ No newline at end of file + return [self.first_stage_model.decode(zu) for zu in z] + +@MODELS.register_class() +class LatentDiffusionFluxMRRedux(LatentDiffusionFluxMR): + para_dict = { + } + para_dict.update(LatentDiffusionFluxMR.para_dict) + + def init_params(self): + super().init_params() + self.redux_adapter_cfg = self.cfg.get("REDUX_ADAPTER", None) + + + def construct_network(self): + super().construct_network() + if self.redux_adapter_cfg is not None: + self.redux_adapter = EMBEDDERS.build(self.redux_adapter_cfg, logger=self.logger).eval().requires_grad_(False) + + def forward_train(self, + image=None, + noise=None, + prompt=[], + **kwargs): + if check_list_of_list(prompt): + prompt = [pp[0] for pp in prompt] + assert self.cond_stage_model is not None + gc_seg = kwargs.pop("gc_seg", []) + gc_seg = int(gc_seg[0]) if len(gc_seg) > 0 else 0 + context = getattr(self.cond_stage_model, 'encode')(prompt) + + image = to_device(image) + x_start = self.encode_first_stage(image, **kwargs) + loss_mask, _ = pack_imagelist_into_tensor(tuple(torch.ones_like(ix, dtype=torch.bool, device=ix.device) for ix in x_start)) + x_start, x_shapes = pack_imagelist_into_tensor(x_start) + context['x_shapes'] = x_shapes + guide_scale = self.guide_scale + if guide_scale is not None: + guide_scale = torch.full((x_start.shape[0],), guide_scale, device=x_start.device, dtype=x_start.dtype) + else: + guide_scale = None + loss = self.diffusion.loss(x_0=x_start, + model=self.model, + model_kwargs={"cond": context, + "gc_seg": gc_seg, + "guidance": guide_scale}, + noise=None, + reduction='none', + **kwargs) + loss = loss[loss_mask].mean() + ret = {'loss': loss, 'probe_data': {'prompt': prompt}} + return ret + + @torch.no_grad() + def forward_sample(self, + noise = None, + prompt=None, + sampler='flow_euler', + sample_steps=20, + guide_scale=3.5, + show_process=True, + x = None, + reverse_scale = -1., + **kwargs + ): + noise, x_shapes = pack_imagelist_into_tensor(noise) + if x is not None: + x, _ = pack_imagelist_into_tensor(x) + context = getattr(self.cond_stage_model, 'encode')(prompt) + context["x_shapes"] = x_shapes + guide_scale = guide_scale or self.guide_scale + if guide_scale is not None: + guide_scale = torch.full((noise.shape[0],), guide_scale, device=noise.device, dtype=noise.dtype) + else: + guide_scale = None + # UNet use input n_prompt + model = self.model_ema if self.use_ema and self.eval_ema else self.model + embedding_context = model.no_sync if isinstance(model, torch.distributed.fsdp.FullyShardedDataParallel) \ + else nullcontext + with embedding_context(): + x_samples = self.diffusion.sample( + noise=noise, + sampler=sampler, + model=self.model, + model_kwargs={"cond": context, "guidance": guide_scale, "gc_seg": -1}, + steps=sample_steps, + show_progress=True, + guide_scale=guide_scale, + return_intermediate=None, + reverse_scale = reverse_scale, + x = x, + **kwargs).float() + x_samples = unpack_tensor_into_imagelist(x_samples, x_shapes) + with torch.autocast(device_type="cuda", dtype=torch.bfloat16): + x_samples = self.decode_first_stage(x_samples) + return x_samples + @torch.no_grad() + def forward_test(self, + image=None, + prompt=[], + sampler='flow_euler', + sample_steps=20, + seed=2023, + guide_scale=3.5, + guide_rescale=0.0, + show_process=True, + log_num = -1, + **kwargs): + + if check_list_of_list(prompt): + prompt = [pp[0] for pp in prompt] + assert self.cond_stage_model is not None + # gc_seg is unused + prompt, image = limit_batch_data([prompt, image], log_num) + seed = seed if seed >= 0 else random.randint(0, 2**32 - 1) + + if 'index' in kwargs: + kwargs.pop('index') + if image is not None: + noise = [self.noise_sample(1, ix.shape[1], ix.shape[2], seed) for ix in image] + else: + 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(1, height, width, seed) for _ in prompt] + + x_samples = self.forward_sample( + prompt=prompt, + sampler=sampler, + sample_steps=sample_steps, + guide_scale=guide_scale, + show_process=show_process, + noise=noise, + ) + + + outputs = list() + for i in range(len(prompt)): + rec_img = torch.clamp((x_samples[i].float() + 1.0) / 2.0, min=0.0, max=1.0) + rec_img = rec_img.squeeze(0) + one_tup = {'prompt': prompt[i], 'n_prompt': '', 'image': rec_img} + outputs.append(one_tup) + return outputs + @staticmethod + def get_config_template(): + return dict_to_yaml('MODEL', + __class__.__name__, + LatentDiffusionFluxMR.para_dict, + set_name=True) diff --git a/scepter/modules/model/network/ldm/ldm_xl.py b/scepter/modules/model/network/ldm/ldm_xl.py index 502ca11..5d20ce1 100644 --- a/scepter/modules/model/network/ldm/ldm_xl.py +++ b/scepter/modules/model/network/ldm/ldm_xl.py @@ -80,7 +80,7 @@ class LatentDiffusionXL(LatentDiffusion): from safetensors.torch import load_file as load_safetensors sd = load_safetensors(path) else: - sd = torch.load(path, map_location='cpu') + sd = torch.load(path, map_location='cpu', weights_only=True) new_sd = OrderedDict() for k, v in sd.items(): ignored = False diff --git a/scepter/modules/model/registry.py b/scepter/modules/model/registry.py index a6032da..9a2e9ed 100644 --- a/scepter/modules/model/registry.py +++ b/scepter/modules/model/registry.py @@ -19,7 +19,7 @@ def build_model(cfg, registry, logger=None, *args, **kwargs): raise TypeError('Pretrain parameter must be a string or list') else: pretrain_cfg = None - + device = cfg.get("DEVICE", None) model = build_from_config(cfg, registry, logger=logger, *args, **kwargs) if pretrain_cfg is not None: if hasattr(model, 'load_pretrained_model'): @@ -49,7 +49,7 @@ def build_diffusion_sampler(cfg, registry, logger=None, *args, **kwargs): MODELS = Registry('MODELS', build_func=build_model) -TOKENIZERS = Registry('TOKENIZER', build_func=build_model) +TOKENIZERS = Registry('TOKENIZERS', build_func=build_model) EMBEDDERS = Registry('EMBEDDERS', build_func=build_model) BACKBONES = Registry('BACKBONES', build_func=build_model) NECKS = Registry('NECKS', build_func=build_model) diff --git a/scepter/modules/model/tokenizer/__init__.py b/scepter/modules/model/tokenizer/__init__.py index a7a9df6..0a8866a 100644 --- a/scepter/modules/model/tokenizer/__init__.py +++ b/scepter/modules/model/tokenizer/__init__.py @@ -1,6 +1,25 @@ # -*- coding: utf-8 -*- # Copyright (c) Alibaba, Inc. and its affiliates. -from scepter.modules.model.tokenizer.base_tokenizer import BaseTokenizer -from scepter.modules.model.tokenizer.tokenizer import (ClipTokenizer, - HuggingfaceTokenizer, - OpenClipTokenizer) +from typing import TYPE_CHECKING +from scepter.modules.utils.import_utils import LazyImportModule + + +if TYPE_CHECKING: + from scepter.modules.model.tokenizer.base_tokenizer import BaseTokenizer + from scepter.modules.model.tokenizer.tokenizer import (ClipTokenizer, + HuggingfaceTokenizer, + OpenClipTokenizer) +else: + _import_structure = { + 'base_tokenizer': ['BaseTokenizer'], + 'tokenizer': ['ClipTokenizer', 'HuggingfaceTokenizer', 'OpenClipTokenizer'] + } + + import sys + sys.modules[__name__] = LazyImportModule( + __name__, + globals()['__file__'], + _import_structure, + module_spec=__spec__, + extra_objects={}, + ) diff --git a/scepter/modules/model/tokenizer/tokenizer_component.py b/scepter/modules/model/tokenizer/tokenizer_component.py index c8be2bb..1467ca0 100644 --- a/scepter/modules/model/tokenizer/tokenizer_component.py +++ b/scepter/modules/model/tokenizer/tokenizer_component.py @@ -152,9 +152,9 @@ def heavy_clean(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 + re.compile(r'[' + '#®•©™&@·º½¾¿¡§~' + r'\)' + r'\(' + r'\]' + # noqa + r'\[' + # noqa + r'\}' + r'\{' + r'\|' + '\\' + r'\/' + r'\*' + # noqa r']{1,}'), # noqa r' ', text) # ***AUSVERKAUFT***, #AUSVERKAUFT diff --git a/scepter/modules/model/tuner/__init__.py b/scepter/modules/model/tuner/__init__.py index bbb2076..741bf72 100644 --- a/scepter/modules/model/tuner/__init__.py +++ b/scepter/modules/model/tuner/__init__.py @@ -1,5 +1,25 @@ # -*- coding: utf-8 -*- # Copyright (c) Alibaba, Inc. and its affiliates. -from scepter.modules.model.tuner import sce -from scepter.modules.model.tuner.swift_tuner import (SwiftPart, SwiftAdapter, SwiftFull, - SwiftLoRA, SwiftSCETuning) +from typing import TYPE_CHECKING +from scepter.modules.utils.import_utils import LazyImportModule + + +if TYPE_CHECKING: + from scepter.modules.model.tuner import sce + from scepter.modules.model.tuner.swift_tuner import (SwiftPart, SwiftAdapter, SwiftFull, + SwiftLoRA, SwiftSCETuning) +else: + _import_structure = { + 'tuner': ['sce'], + 'swift_tuner': ['SwiftPart', 'SwiftAdapter', 'SwiftFull', + 'SwiftLoRA', 'SwiftSCETuning'] + } + + import sys + sys.modules[__name__] = LazyImportModule( + __name__, + globals()['__file__'], + _import_structure, + module_spec=__spec__, + extra_objects={}, + ) diff --git a/scepter/modules/model/tuner/sce/__init__.py b/scepter/modules/model/tuner/sce/__init__.py index 8c549fe..7f5bf2a 100644 --- a/scepter/modules/model/tuner/sce/__init__.py +++ b/scepter/modules/model/tuner/sce/__init__.py @@ -1,4 +1,23 @@ # -*- coding: utf-8 -*- # Copyright (c) Alibaba, Inc. and its affiliates. -from scepter.modules.model.tuner.sce.scetuning import CSCTuners, SCTuner -from scepter.modules.model.tuner.sce.scetuning_component import SCEAdapter +from typing import TYPE_CHECKING +from scepter.modules.utils.import_utils import LazyImportModule + + +if TYPE_CHECKING: + from scepter.modules.model.tuner.sce.scetuning import CSCTuners, SCTuner + from scepter.modules.model.tuner.sce.scetuning_component import SCEAdapter +else: + _import_structure = { + 'scetuning': ['CSCTuners', 'SCTuner'], + 'scetuning_component': ['SCEAdapter'] + } + + import sys + sys.modules[__name__] = LazyImportModule( + __name__, + globals()['__file__'], + _import_structure, + module_spec=__spec__, + extra_objects={}, + ) diff --git a/scepter/modules/model/tuner/sce/scetuning.py b/scepter/modules/model/tuner/sce/scetuning.py index 827b15d..7b68c99 100644 --- a/scepter/modules/model/tuner/sce/scetuning.py +++ b/scepter/modules/model/tuner/sce/scetuning.py @@ -153,7 +153,7 @@ class CSCTuners(BaseTuner): def init_from_ckpt(self, path): model_new = OrderedDict() - model = torch.load(path, map_location='cpu') + model = torch.load(path, map_location='cpu', weights_only=True) for k, v in model.items(): if k.startswith('model.'): k = k[len('model.'):] diff --git a/scepter/modules/model/tuner/swift_tuner.py b/scepter/modules/model/tuner/swift_tuner.py index 68a195f..8ddf7ad 100644 --- a/scepter/modules/model/tuner/swift_tuner.py +++ b/scepter/modules/model/tuner/swift_tuner.py @@ -79,6 +79,7 @@ class SwiftLoRA(): lora_alpha=cfg.LORA_ALPHA, lora_dropout=cfg.LORA_DROPOUT, bias=cfg.BIAS, + use_dora=cfg.get('USE_DORA', False), target_modules=cfg.TARGET_MODULES) def __call__(self, *args, **kwargs): diff --git a/scepter/modules/opt/__init__.py b/scepter/modules/opt/__init__.py index 7779fc5..7255046 100644 --- a/scepter/modules/opt/__init__.py +++ b/scepter/modules/opt/__init__.py @@ -1,4 +1,21 @@ # -*- coding: utf-8 -*- # Copyright (c) Alibaba, Inc. and its affiliates. +from typing import TYPE_CHECKING +from scepter.modules.utils.import_utils import LazyImportModule -from scepter.modules.opt import lr_schedulers, optimizers + +if TYPE_CHECKING: + from scepter.modules.opt import lr_schedulers, optimizers +else: + _import_structure = { + 'opt': ['lr_schedulers', 'optimizers'] + } + + import sys + sys.modules[__name__] = LazyImportModule( + __name__, + globals()['__file__'], + _import_structure, + module_spec=__spec__, + extra_objects={}, + ) diff --git a/scepter/modules/opt/lr_schedulers/__init__.py b/scepter/modules/opt/lr_schedulers/__init__.py index 12b9340..2600d3d 100644 --- a/scepter/modules/opt/lr_schedulers/__init__.py +++ b/scepter/modules/opt/lr_schedulers/__init__.py @@ -1,7 +1,29 @@ # -*- coding: utf-8 -*- # Copyright (c) Alibaba, Inc. and its affiliates. +from typing import TYPE_CHECKING +from scepter.modules.utils.import_utils import LazyImportModule -from scepter.modules.opt.lr_schedulers.define_schedulers import LinoPolyLR -from scepter.modules.opt.lr_schedulers.official_schedulers import * # noqa -from scepter.modules.opt.lr_schedulers.warmup import (StepAnnealingLR, - WarmupToConstantLR) + +if TYPE_CHECKING: + from scepter.modules.opt.lr_schedulers.define_schedulers import LinoPolyLR + from scepter.modules.opt.lr_schedulers.official_schedulers import * # noqa + from scepter.modules.opt.lr_schedulers.warmup import (StepAnnealingLR, + WarmupToConstantLR) +else: + _import_structure = { + 'define_schedulers': ['LinoPolyLR'], + 'official_schedulers': ['StepLR', 'CyclicLR', 'LambdaLR', 'MultiStepLR', + 'ExponentialLR', 'CosineAnnealingLR', + 'CosineAnnealingWarmRestarts', 'ReduceLROnPlateau'], + 'warmup': ['StepAnnealingLR', 'WarmupToConstantLR'], + 'registry': ['LR_SCHEDULERS'] + } + + import sys + sys.modules[__name__] = LazyImportModule( + __name__, + globals()['__file__'], + _import_structure, + module_spec=__spec__, + extra_objects={}, + ) diff --git a/scepter/modules/opt/lr_schedulers/registry.py b/scepter/modules/opt/lr_schedulers/registry.py index e1bc28a..0f321f5 100644 --- a/scepter/modules/opt/lr_schedulers/registry.py +++ b/scepter/modules/opt/lr_schedulers/registry.py @@ -18,8 +18,12 @@ def build_lr_scheduler(cfg, registry, logger=None, *args, **kwargs): cfg = deep_copy(cfg) assert kwargs is not None and 'optimizer' in kwargs optimizer = kwargs['optimizer'] - req_type = cfg.get('NAME') + + from scepter.modules.utils.import_utils import LazyImportModule + sig = (registry.name.upper(), req_type) + LazyImportModule.import_module(sig) + if isinstance(req_type, str): req_type_entry = registry.get(req_type) if req_type_entry is None: diff --git a/scepter/modules/opt/optimizers/__init__.py b/scepter/modules/opt/optimizers/__init__.py index 675bcbb..dc91328 100644 --- a/scepter/modules/opt/optimizers/__init__.py +++ b/scepter/modules/opt/optimizers/__init__.py @@ -1,7 +1,27 @@ # -*- coding: utf-8 -*- # Copyright (c) Alibaba, Inc. and its affiliates. +from typing import TYPE_CHECKING +from scepter.modules.utils.import_utils import LazyImportModule -from scepter.modules.opt.optimizers.official_optimizers import ( - ASGD, LBFGS, SGD, Adadelta, Adagrad, Adam, Adamax, AdamW, RMSprop, Rprop, - SparseAdam) -from scepter.modules.opt.optimizers.registry import OPTIMIZERS + +if TYPE_CHECKING: + from scepter.modules.opt.optimizers.official_optimizers import ( + ASGD, LBFGS, SGD, Adadelta, Adagrad, Adam, Adamax, AdamW, RMSprop, Rprop, + SparseAdam) + from scepter.modules.opt.optimizers.registry import OPTIMIZERS +else: + _import_structure = { + 'official_optimizers': ['ASGD', 'LBFGS', 'SGD', 'Adadelta', + 'Adagrad', 'Adam', 'Adamax', 'AdamW', + 'RMSprop', 'Rprop', 'SparseAdam'], + 'registry': ['OPTIMIZERS'] + } + + import sys + sys.modules[__name__] = LazyImportModule( + __name__, + globals()['__file__'], + _import_structure, + module_spec=__spec__, + extra_objects={}, + ) diff --git a/scepter/modules/opt/optimizers/registry.py b/scepter/modules/opt/optimizers/registry.py index b3bc53a..b5da424 100644 --- a/scepter/modules/opt/optimizers/registry.py +++ b/scepter/modules/opt/optimizers/registry.py @@ -19,8 +19,12 @@ def build_optimizer(cfg, registry, logger=None, *args, **kwargs): parameters = kwargs['parameters'] cfg = deep_copy(cfg) - req_type = cfg.get('NAME') + + from scepter.modules.utils.import_utils import LazyImportModule + sig = (registry.name.upper(), req_type) + LazyImportModule.import_module(sig) + if isinstance(req_type, str): req_type_entry = registry.get(req_type) if req_type_entry is None: diff --git a/scepter/modules/solver/__init__.py b/scepter/modules/solver/__init__.py index 7b0a9a7..e4d1ee5 100644 --- a/scepter/modules/solver/__init__.py +++ b/scepter/modules/solver/__init__.py @@ -1,8 +1,33 @@ # -*- coding: utf-8 -*- # Copyright (c) Alibaba, Inc. and its affiliates. -from scepter.modules.solver import hooks -from scepter.modules.solver.base_solver import BaseSolver -from scepter.modules.solver.diffusion_solver import LatentDiffusionSolver -from scepter.modules.solver.train_val_solver import TrainValSolver -from scepter.modules.solver.ace_solver import ACESolver -from scepter.modules.solver.diffusion_video_solver import LatentDiffusionVideoSolver \ No newline at end of file +from typing import TYPE_CHECKING +from scepter.modules.utils.import_utils import LazyImportModule + + +if TYPE_CHECKING: + from scepter.modules.solver import hooks + from scepter.modules.solver.base_solver import BaseSolver + from scepter.modules.solver.diffusion_solver import LatentDiffusionSolver + from scepter.modules.solver.train_val_solver import TrainValSolver + from scepter.modules.solver.ace_solver import ACESolver + from scepter.modules.solver.ace_plus_solver import ACEPlusSolver + from scepter.modules.solver.diffusion_video_solver import LatentDiffusionVideoSolver +else: + _import_structure = { + 'solver': ['hooks'], + 'base_solver': ['BaseSolver'], + 'diffusion_solver': ['LatentDiffusionSolver'], + 'train_val_solver': ['TrainValSolver'], + 'ace_solver': ['ACESolver'], + 'ace_plus_solver': ['ACEPlusSolver'], + 'diffusion_video_solver': ['LatentDiffusionVideoSolver'] + } + + import sys + sys.modules[__name__] = LazyImportModule( + __name__, + globals()['__file__'], + _import_structure, + module_spec=__spec__, + extra_objects={}, + ) diff --git a/scepter/modules/solver/ace_plus_solver.py b/scepter/modules/solver/ace_plus_solver.py new file mode 100644 index 0000000..90d3fc8 --- /dev/null +++ b/scepter/modules/solver/ace_plus_solver.py @@ -0,0 +1,164 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +import numpy as np +import torch +from scepter.modules.solver import LatentDiffusionSolver +from scepter.modules.solver.registry import SOLVERS +from scepter.modules.utils.data import transfer_data_to_cuda +from scepter.modules.utils.distribute import we +from scepter.modules.utils.probe import ProbeData +from tqdm import tqdm +@SOLVERS.register_class() +class ACEPlusSolver(LatentDiffusionSolver): + def __init__(self, cfg, logger=None): + super().__init__(cfg, logger=logger) + self.probe_prompt = cfg.get("PROBE_PROMPT", None) + self.probe_hw = cfg.get("PROBE_HW", []) + @torch.no_grad() + def run_eval(self): + self.eval_mode() + self.before_all_iter(self.hooks_dict[self._mode]) + all_results = [] + for batch_idx, batch_data in tqdm( + enumerate(self.datas[self._mode].dataloader)): + self.before_iter(self.hooks_dict[self._mode]) + 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): + results = self.run_step_eval(transfer_data_to_cuda(batch_data), + batch_idx, + step=self.total_iter, + rank=we.rank) + all_results.extend(results) + self.after_iter(self.hooks_dict[self._mode]) + log_data, log_label = self.save_results(all_results) + self.register_probe({'eval_label': log_label}) + self.register_probe({ + 'eval_image': + ProbeData(log_data, + is_image=True, + build_html=True, + build_label=log_label) + }) + self.after_all_iter(self.hooks_dict[self._mode]) + + @torch.no_grad() + def run_test(self): + self.test_mode() + self.before_all_iter(self.hooks_dict[self._mode]) + all_results = [] + for batch_idx, batch_data in tqdm( + enumerate(self.datas[self._mode].dataloader)): + self.before_iter(self.hooks_dict[self._mode]) + 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): + results = self.run_step_eval(transfer_data_to_cuda(batch_data), + batch_idx, + step=self.total_iter, + rank=we.rank) + all_results.extend(results) + self.after_iter(self.hooks_dict[self._mode]) + log_data, log_label = self.save_results(all_results) + self.register_probe({'test_label': log_label}) + self.register_probe({ + 'test_image': + ProbeData(log_data, + is_image=True, + build_html=True, + build_label=log_label) + }) + + self.after_all_iter(self.hooks_dict[self._mode]) + + def save_results(self, results): + log_data, log_label = [], [] + for result in results: + ret_images, ret_labels = [], [] + edit_image = result.get('edit_image', None) + edit_mask = result.get('edit_mask', None) + if edit_image is not None: + for i, edit_img in enumerate(result['edit_image']): + if edit_img is None: + continue + ret_images.append((edit_img.permute(1, 2, 0).cpu().numpy() * 255).astype(np.uint8)) + ret_labels.append(f'edit_image{i}; ') + if edit_mask is not None: + ret_images.append((edit_mask[i].permute(1, 2, 0).cpu().numpy() * 255).astype(np.uint8)) + ret_labels.append(f'edit_mask{i}; ') + + target_image = result.get('target_image', None) + target_mask = result.get('target_mask', None) + if target_image is not None: + ret_images.append((target_image.permute(1, 2, 0).cpu().numpy() * 255).astype(np.uint8)) + ret_labels.append(f'target_image; ') + if target_mask is not None: + ret_images.append((target_mask.permute(1, 2, 0).cpu().numpy() * 255).astype(np.uint8)) + ret_labels.append(f'target_mask; ') + teacher_image = result.get('image', None) + if teacher_image is not None: + ret_images.append((teacher_image.permute(1, 2, 0).cpu().numpy() * 255).astype(np.uint8)) + ret_labels.append(f"teacher_image") + reconstruct_image = result.get('reconstruct_image', None) + if reconstruct_image is not None: + ret_images.append((reconstruct_image.permute(1, 2, 0).cpu().numpy() * 255).astype(np.uint8)) + ret_labels.append(f"{result['instruction']}") + log_data.append(ret_images) + log_label.append(ret_labels) + return log_data, log_label + @property + def probe_data(self): + if not we.debug and self.mode == 'train': + batch_data = transfer_data_to_cuda(self.current_batch_data[self.mode]) + self.eval_mode() + with torch.autocast(device_type='cuda', + enabled=self.use_amp, + dtype=self.dtype): + batch_data['log_num'] = self.log_train_num + batch_data.update(self.sample_args.get_lowercase_dict()) + results = self.run_step_eval(batch_data) + self.train_mode() + log_data, log_label = self.save_results(results) + self.register_probe({ + 'train_image': + ProbeData(log_data, + is_image=True, + build_html=True, + build_label=log_label) + }) + self.register_probe({'train_label': log_label}) + if self.probe_prompt: + self.eval_mode() + all_results = [] + for prompt in self.probe_prompt: + with torch.autocast(device_type='cuda', + enabled=self.use_amp, + dtype=self.dtype): + batch_data = { + "prompt": [[prompt]], + "image": [torch.zeros(3, self.probe_hw[0], self.probe_hw[1])], + "image_mask": [torch.ones(1, self.probe_hw[0], self.probe_hw[1])], + "src_image_list": [[]], + "src_mask_list": [[]], + "edit_id": [[]], + "height": self.probe_hw[0], + "width": self.probe_hw[1] + } + batch_data.update(self.sample_args.get_lowercase_dict()) + results = self.run_step_eval(batch_data) + all_results.extend(results) + self.train_mode() + log_data, log_label = self.save_results(all_results) + self.register_probe({ + 'probe_image': + ProbeData(log_data, + is_image=True, + build_html=True, + build_label=log_label) + }) + + return super(LatentDiffusionSolver, self).probe_data diff --git a/scepter/modules/solver/diffusion_solver.py b/scepter/modules/solver/diffusion_solver.py index 2d90d7b..be2294c 100644 --- a/scepter/modules/solver/diffusion_solver.py +++ b/scepter/modules/solver/diffusion_solver.py @@ -217,6 +217,7 @@ class LatentDiffusionSolver(BaseSolver): self.tuner_cfg = cfg.get('TUNER', None) self.freeze_cfg = cfg.get('FREEZE', None) self.log_train_num = cfg.get("LOG_TRAIN_NUM", -1) + self.timesteps = cfg.get("TIMESTEPS", 1000) def set_up(self): self.construct_data() @@ -272,6 +273,12 @@ class LatentDiffusionSolver(BaseSolver): module_keys = [key for key, _ in self.model.named_modules()] self.logger.info(module_keys) + def train_parameters(self): + model = self.model + for key, val in model.named_parameters(): + if val.requires_grad: + yield val + def model_to_device(self): self.model = self.model.to(we.device_id) @@ -285,8 +292,15 @@ class LatentDiffusionSolver(BaseSolver): self.cfg.OPTIMIZER.LEARNING_RATE *= all_batch_size self.cfg.OPTIMIZER.LEARNING_RATE /= 640 + def get_params(self, module): + train_params = [] + for param in module.parameters(): + if param.requires_grad: + train_params.append(param) + return train_params + def init_opti(self): - import torch.cuda.amp as amp + import torch.amp as amp import torch.distributed as dist if we.is_distributed: @@ -383,7 +397,7 @@ class LatentDiffusionSolver(BaseSolver): for module in self.train_modules: if hasattr(self.model, module): current_module = getattr(self.model, module) - train_params += list(current_module.parameters()) + train_params += self.get_params(current_module) self.optimizer = OPTIMIZERS.build(self.cfg.OPTIMIZER, logger=self.logger, @@ -398,12 +412,12 @@ class LatentDiffusionSolver(BaseSolver): self.optimizer = OPTIMIZERS.build( self.cfg.OPTIMIZER, logger=self.logger, - parameters=self.model.parameters()) + parameters=self.get_params(self.model)) else: self.optimizer = OPTIMIZERS.build( self.cfg.OPTIMIZER, logger=self.logger, - parameters=self.model.parameters()) + parameters=self.get_params(self.model)) if self.cfg.have('LR_SCHEDULER') and self.optimizer is not None: self.cfg.LR_SCHEDULER.TOTAL_STEPS = self.max_steps @@ -422,8 +436,8 @@ class LatentDiffusionSolver(BaseSolver): process_group=None) else: self.scaler = amp.GradScaler(enabled=self.enable_gradscaler) - elif self.cfg.DTYPE in ['float16']: - self.scaler = amp.GradScaler() + elif self.cfg.DTYPE in ['float16', 'bfloat16']: + self.scaler = amp.GradScaler(enabled=self.enable_gradscaler) else: self.scaler = None else: @@ -736,16 +750,40 @@ class LatentDiffusionSolver(BaseSolver): if model is None: model = self.model - swift_cfg_dict = {} - for t_id, t_cfg in enumerate(tuner_cfg): - cfg_name = t_cfg['NAME'] - init_config = TUNERS.build(t_cfg, logger=self.logger)() - if init_config is None: - continue - swift_cfg_dict[f'{t_id}_{cfg_name}'] = init_config - if len(swift_cfg_dict) > 0: + + if isinstance(tuner_cfg, str): from swift import Swift - model = Swift.prepare_model(self.model, config=swift_cfg_dict, autocast_adapter_dtype=False) + from scepter.modules.utils.file_system import FS + with FS.get_dir_to_local_dir(tuner_cfg, wait_finish=True) as local_dir: + model = Swift.from_pretrained(model, local_dir, autocast_adapter_dtype=False) + self.logger.info(f'Load tuner model from {tuner_cfg}') + else: + swift_cfg_dict = {} + swfit_ckpts = {} + for t_id, t_cfg in enumerate(tuner_cfg): + if 'PRETRAINED_MODEL' in t_cfg: + pretrained_model = t_cfg.pop('PRETRAINED_MODEL') + from scepter.modules.utils.file_system import FS + with FS.get_from(pretrained_model, wait_finish=True) as local_path: + if local_path.endswith('safetensors'): + from safetensors.torch import load_file as load_safetensors + ckpt = load_safetensors(local_path) + else: + ckpt = torch.load(local_path, map_location='cpu', weights_only=True) + swfit_ckpts.update(ckpt) + cfg_name = t_cfg['NAME'] + init_config = TUNERS.build(t_cfg, logger=self.logger)() + if init_config is None: + continue + swift_cfg_dict[f'{t_id}_{cfg_name}'] = init_config + + if len(swift_cfg_dict) > 0: + from swift import Swift + model = Swift.prepare_model(self.model, config=swift_cfg_dict, autocast_adapter_dtype=False) + if len(swfit_ckpts) > 0: + swfit_ckpts = {k.replace('transformer.', 'model.').replace('lora_A.weight', 'lora_A.0_SwiftLoRA.weight').replace('lora_B.weight', 'lora_B.0_SwiftLoRA.weight'): v for k, v in swfit_ckpts.items()} + model.load_state_dict(swfit_ckpts, strict=True) + self.logger.info(f'Restored from TUNER with length of {len(swfit_ckpts)}') self.logger.info([(key, param.shape) for key, param in model.named_parameters() if param.requires_grad]) return model diff --git a/scepter/modules/solver/diffusion_video_solver.py b/scepter/modules/solver/diffusion_video_solver.py index 3892d92..9ca4c22 100644 --- a/scepter/modules/solver/diffusion_video_solver.py +++ b/scepter/modules/solver/diffusion_video_solver.py @@ -24,23 +24,31 @@ class LatentDiffusionVideoSolver(LatentDiffusionSolver): for result in results: ret_videos, ret_labels = [], [] if 'edit_video' in result: - ret_videos.append((result['edit_video'].permute(1, 2, 3, 0).cpu().numpy() * - 255).astype(np.uint8)) + ret_videos.append((result['edit_video'].permute(1, 2, 3, 0)*255).cpu().numpy().astype(np.uint8)) ret_labels.append("left: edit video") if 'edit_image' in result: - ret_videos.append((result['edit_image'].permute(1, 2, 3, 0).cpu().numpy() * - 255).astype(np.uint8)) + ret_videos.append((result['edit_image'].permute(1, 2, 3, 0)*255).cpu().numpy().astype(np.uint8)) ret_labels.append("left: edit image") + if 'edit_mask' in result: + if len(result['edit_mask'].shape) == 4: + ret_videos.append((result['edit_mask'].permute(1, 2, 3, 0)*255).cpu().numpy().astype(np.uint8)) + elif len(result['edit_mask'].shape) == 3: + if result['edit_mask'].shape[0] == 1: + result['edit_mask'] = result['edit_mask'].repeat(3, 1, 1) + ret_videos.append(((result['edit_mask'].permute(1, 2, 0)*255).cpu().numpy()[None, ...]).astype(np.uint8)) + else: + if result['edit_mask'].shape[0] == 1: + result['edit_mask'] = result['edit_mask'].repeat(3, 1, 1, 1) + ret_videos.append(((result['edit_mask'].permute(1, 2, 3, 0)*255).cpu().numpy()).astype(np.uint8)) + ret_labels.append("middle: edit mask") if 'target_video' in result: if len(ret_videos) > 0: ret_labels.append("middle: target video") else: ret_labels.append("left: target video") - ret_videos.append((result['target_video'].permute(1, 2, 3, 0).cpu().numpy() * - 255).astype(np.uint8)) + ret_videos.append((result['target_video'].permute(1, 2, 3, 0)*255).cpu().numpy().astype(np.uint8)) - ret_videos.append((result['reconstruct_video'].permute(1, 2, 3, 0).cpu().numpy() * - 255).astype(np.uint8)) + ret_videos.append((result['reconstruct_video'].permute(1, 2, 3, 0)*255).cpu().numpy().astype(np.uint8)) ret_labels.append("right: generation video" + " Prompt: " + result['instruction']) log_data.append(ret_videos) @@ -70,15 +78,11 @@ class LatentDiffusionVideoSolver(LatentDiffusionSolver): 'batch_size': len(batch_data['prompt']) }) self.current_batch_data[self.mode] = batch_data - if self.sample_args: - self.current_batch_data[self.mode].update( - self.sample_args.get_lowercase_dict()) - batch_data = transfer_data_to_cuda(batch_data) with torch.autocast(device_type='cuda', enabled=self.use_amp, dtype=self.dtype): results = self.run_step_train( - batch_data, + transfer_data_to_cuda(batch_data), step, step=self.total_iter, rank=we.rank) @@ -124,6 +128,21 @@ class LatentDiffusionVideoSolver(LatentDiffusionSolver): }) self.after_all_iter(self.hooks_dict[self._mode]) + def run_step_val(self, batch_data, noise_generator=None): + loss_dict = {} + batch_data = transfer_data_to_cuda(batch_data) + with torch.autocast(device_type='cuda', + enabled=self.use_amp, + dtype=self.dtype): + if hasattr(self.model, 'module'): + results = self.model.module.forward_train(**batch_data) + else: + results = self.model.forward_train(**batch_data) + loss = results['loss'] + for sample_id in batch_data['sample_id']: + loss_dict[sample_id] = loss.detach().cpu().numpy() + return loss_dict + @torch.no_grad() def run_test(self): self.test_mode() @@ -166,7 +185,7 @@ class LatentDiffusionVideoSolver(LatentDiffusionSolver): with torch.autocast(device_type='cuda', enabled=self.use_amp, dtype=self.dtype): - batch_data['log_train_num'] = self.log_train_num + batch_data['log_num'] = self.log_train_num all_results = self.run_step_eval(transfer_data_to_cuda(batch_data)) self.train_mode() log_data, log_label = self.save_results(all_results) diff --git a/scepter/modules/solver/hooks/__init__.py b/scepter/modules/solver/hooks/__init__.py index 1bd393f..b0af944 100644 --- a/scepter/modules/solver/hooks/__init__.py +++ b/scepter/modules/solver/hooks/__init__.py @@ -1,16 +1,7 @@ # -*- coding: utf-8 -*- # Copyright (c) Alibaba, Inc. and its affiliates. - -from scepter.modules.solver.hooks.backward import BackwardHook -from scepter.modules.solver.hooks.checkpoint import CheckpointHook -from scepter.modules.solver.hooks.data_probe import ProbeDataHook -from scepter.modules.solver.hooks.ema import ModelEmaHook -from scepter.modules.solver.hooks.hook import Hook -from scepter.modules.solver.hooks.log import LogHook, TensorboardLogHook -from scepter.modules.solver.hooks.lr import LrHook -from scepter.modules.solver.hooks.registry import HOOKS -from scepter.modules.solver.hooks.safetensors import SafetensorsHook -from scepter.modules.solver.hooks.sampler import DistSamplerHook +from typing import TYPE_CHECKING +from scepter.modules.utils.import_utils import LazyImportModule """ Normally, hooks have priorities, below we recommend priority that runs fine (low score MEANS high priority) BackwardHook: 0 @@ -46,8 +37,39 @@ after solve: TensorboardLogHook: close file handler """ -__all__ = [ - 'HOOKS', 'BackwardHook', 'CheckpointHook', 'Hook', 'LrHook', 'LogHook', - 'TensorboardLogHook', 'DistSamplerHook', 'ProbeDataHook', - 'SafetensorsHook', 'ModelEmaHook' -] + +if TYPE_CHECKING: + from scepter.modules.solver.hooks.backward import BackwardHook + from scepter.modules.solver.hooks.checkpoint import CheckpointHook + from scepter.modules.solver.hooks.data_probe import ProbeDataHook + from scepter.modules.solver.hooks.ema import ModelEmaHook + from scepter.modules.solver.hooks.hook import Hook + from scepter.modules.solver.hooks.log import LogHook, TensorboardLogHook + from scepter.modules.solver.hooks.lr import LrHook + from scepter.modules.solver.hooks.registry import HOOKS + from scepter.modules.solver.hooks.safetensors import SafetensorsHook + from scepter.modules.solver.hooks.sampler import DistSamplerHook + from scepter.modules.solver.hooks.val_loss import ValLossHook +else: + _import_structure = { + 'backward': ['BackwardHook'], + 'checkpoint': ['CheckpointHook'], + 'data_probe': ['ProbeDataHook'], + 'ema': ['ModelEmaHook'], + 'hook': ['Hook'], + 'log': ['LogHook', 'TensorboardLogHook'], + 'lr': ['LrHook'], + 'registry': ['HOOKS'], + 'safetensors': ['SafetensorsHook'], + 'sampler': ['DistSamplerHook'], + 'val_loss': ['ValLossHook'] + } + + import sys + sys.modules[__name__] = LazyImportModule( + __name__, + globals()['__file__'], + _import_structure, + module_spec=__spec__, + extra_objects={}, + ) diff --git a/scepter/modules/solver/hooks/backward.py b/scepter/modules/solver/hooks/backward.py index 5dfb38d..3ab2574 100644 --- a/scepter/modules/solver/hooks/backward.py +++ b/scepter/modules/solver/hooks/backward.py @@ -112,10 +112,15 @@ class BackwardHook(Hook): f'Profiler stop after {self.profile_step} steps') FS.put_dir_from_local_dir(self._local_log_dir, self.log_dir) - def grad_clip(self, parameters): - torch.nn.utils.clip_grad_norm_(parameters=parameters, - max_norm=self.gradient_clip, - norm_type=2) + def grad_clip(self, optimizer): + for params_group in optimizer.param_groups: + train_params = [] + for param in params_group['params']: + if param.requires_grad: + train_params.append(param) + # print(len(train_params), self.gradient_clip) + torch.nn.utils.clip_grad_norm_(parameters=train_params, + max_norm=self.gradient_clip) def after_iter(self, solver): if solver.optimizer is not None and solver.is_train_mode: @@ -131,9 +136,9 @@ class BackwardHook(Hook): # Suppose profiler run after backward, so we need to set backward_prev_step # as the previous one step before the backward step if self.current_step % self.accumulate_step == 0: + solver.scaler.unscale_(solver.optimizer) if self.gradient_clip > 0: - solver.scaler.unscale_(solver.optimizer) - self.grad_clip(solver.train_parameters()) + self.grad_clip(solver.optimizer) self.profile(solver) solver.scaler.step(solver.optimizer) solver.scaler.update() @@ -145,7 +150,7 @@ class BackwardHook(Hook): # as the previous one step before the backward step if self.current_step % self.accumulate_step == 0: if self.gradient_clip > 0: - self.grad_clip(solver.train_parameters()) + self.grad_clip(solver.optimizer) self.profile(solver) solver.optimizer.step() solver.optimizer.zero_grad() diff --git a/scepter/modules/solver/hooks/checkpoint.py b/scepter/modules/solver/hooks/checkpoint.py index 61cb071..f0a5a29 100644 --- a/scepter/modules/solver/hooks/checkpoint.py +++ b/scepter/modules/solver/hooks/checkpoint.py @@ -96,7 +96,7 @@ class CheckpointHook(Hook): with FS.get_from(solver.resume_from, wait_finish=True) as local_file: solver.logger.info(f'Loading checkpoint from {solver.resume_from}') checkpoint = torch.load(local_file, - map_location=torch.device('cpu')) + map_location=torch.device('cpu'), weights_only=True) solver.load_checkpoint(checkpoint) if self.save_best and '_CheckpointHook_best' in checkpoint: diff --git a/scepter/modules/solver/hooks/val_loss.py b/scepter/modules/solver/hooks/val_loss.py new file mode 100644 index 0000000..d32cbe5 --- /dev/null +++ b/scepter/modules/solver/hooks/val_loss.py @@ -0,0 +1,230 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +import json +import os + +import numpy as np +import torch +from tqdm import tqdm + +from scepter.modules.data.dataset import DATASETS +from scepter.modules.solver.hooks.hook import Hook +from scepter.modules.solver.hooks.registry import HOOKS +from scepter.modules.utils.config import dict_to_yaml +from scepter.modules.utils.distribute import barrier, gather_data, we +from scepter.modules.utils.file_system import FS +from scepter.modules.utils.math_plot import plot_multi_curves + +_DEFAULT_VAL_PRIORITY = 200 + + +def float_format(o): + if isinstance(o, float): + return f"{o: .6f}" + raise TypeError(f"Type {type(o)} not serializable") + + +@HOOKS.register_class() +class ValLossHook(Hook): + para_dict = [{ + 'PRIORITY': { + 'value': _DEFAULT_VAL_PRIORITY, + 'description': 'The priority for processing!' + }, + 'VAL_INTERVAL': { + 'value': 1000, + 'description': 'the interval for log print!' + }, + 'VAL_LIMITATION_SIZE': { + 'value': 1000000, + 'description': 'the limitation size for validation!' + }, + 'VAL_SEED': { + 'value': 2025, + 'description': 'the validation seed for t or generator sample!' + } + }] + + def __init__(self, cfg, logger=None): + super(ValLossHook, self).__init__(cfg, logger=logger) + self.priority = cfg.get('PRIORITY', _DEFAULT_VAL_PRIORITY) + self.val_interval = cfg.get('VAL_INTERVAL', 1000) + self.val_dim = cfg.get('VAL_DIM', 'all') + self.meta_field = cfg.get('META_FIELD', ['edit_type', 'data_type']) + self.save_folder = cfg.get('SAVE_FOLDER', 'val_loss') + self.val_limitation_size = cfg.get('VAL_LIMITATION_SIZE', 1000000) + self.val_seed = cfg.get('VAL_SEED', 2025) + self.data = DATASETS.build(cfg.DATA, logger=logger) + + def before_all_iter(self, solver): + solver.eval_mode() + self.eval_set_size = len(self.data.dataset) + if self.eval_set_size > self.val_limitation_size: + self.logger.info( + f"The samples number {self.eval_set_size} of validation set " + f"should not great than {self.val_limitation_size}") + assert self.eval_set_size < self.val_limitation_size + if not hasattr(solver, 'run_step_val'): + self.logger.info( + f"The val-loss hook should have the function run_step_val" # noqa + ) # noqa + assert hasattr(solver, 'run_step_val') + if not self.data.batch_size == 1: + self.logger.info( + f"The batch_size of validation set should be 1 " # noqa + f"when you use the validation hook to make the results deterministic." # noqa + ) + assert self.data.batch_size == 1 + timestamp_generator = torch.Generator(device=we.device_id) + timestamp_generator.manual_seed(self.val_seed) + u = torch.rand((self.eval_set_size, ), + device=we.device_id, + generator=timestamp_generator) + self.t = (u * (solver.timesteps - 1)).round().long() + solver.val_interval = self.val_interval + solver.train_mode() + + def get_val_loss(self, solver, step): + all_loss = [] + # batch-size must be 1 + for batch_data in tqdm(self.data.dataloader): + # generate t list + sample_id = int(batch_data['sample_id'][0]) + meta_info = {m_f: batch_data[m_f][0] for m_f in self.meta_field} + meta_info['sample_id'] = sample_id + + batch_data['t'] = torch.stack( + [self.t[sample_id % self.eval_set_size]]) + noise_generator = torch.Generator(device=we.device_id) + noise_generator.manual_seed(sample_id + 10000 * self.val_seed) + # get generator according to the sample_id + with torch.no_grad(): + loss = solver.run_step_val(batch_data, noise_generator) + meta_info['loss'] = float(loss[sample_id]) + all_loss.append(meta_info) + all_loss = json.dumps(all_loss, default=float_format) + all_loss = gather_data([all_loss]) + if we.rank == 0: + reduce_loss = [] + for loss in all_loss: + reduce_loss.extend(json.loads(loss)) + compute_results = self.compute_avg_loss(reduce_loss) + self.save_record(solver, compute_results, reduce_loss, step) + return + + def compute_avg_loss(self, loss_list): + all_avg_ls = [] + avg_ls = {} + for ls in loss_list: + for m_f in self.meta_field: + m_f_v = ls[m_f] + ls_key = m_f + '_' + m_f_v + if ls_key not in avg_ls: + avg_ls[ls_key] = [] + avg_ls[ls_key].append(ls['loss']) + all_avg_ls.append(ls['loss']) + compute_results = { + 'all': sum(all_avg_ls) / len(all_avg_ls), + } + compute_results.update( + {m_f: sum(avg_ls[m_f]) / len(avg_ls[m_f]) + for m_f in avg_ls}) + return compute_results + + def save_record(self, solver, compute_results, all_loss, step): + save_folder = os.path.join(solver.work_dir, self.save_folder) + # save history + save_history = os.path.join(save_folder, 'history.json') + + draw_curve = False + + if FS.exists(save_history): + results = json.loads(FS.get_object(save_history).decode()) + all_loss = {loss['sample_id']: loss for loss in all_loss} + for loss in results['detail']: + loss['loss'] = {int(k): v for k, v in loss['loss'].items()} + loss['loss'][step] = all_loss[loss['sample_id']]['loss'] + for k, v in compute_results.items(): + results['summary'][k] = { + int(kk): vv + for kk, vv in results['summary'][k].items() + } + results['summary'][k][step] = v + draw_curve = True + else: + results = {'detail': [], 'summary': {}} + for loss in all_loss: + loss_v = loss.pop('loss') + loss['loss'] = {step: loss_v} + results['detail'].append(loss) + for k, v in compute_results.items(): + if k not in results['summary']: + results['summary'][k] = {} + results['summary'][k][step] = v + # + FS.put_object( + json.dumps(results, default=float_format).encode(), save_history) + # plot current curve + if draw_curve: + self.plot_results(results['summary'], + os.path.join(save_folder, 'curve')) + # print current log + print_msg = '' + for k, v in compute_results.items(): + print_msg += f"{k}: {v: .4f} " + self.logger.info(f"Step {step} validation loss: {print_msg}") + + def plot_results(self, plot_data, save_folder): + y = [] + steps = [] + # one image + for label, curve_data in plot_data.items(): + curve_data = [[step, value] for step, value in curve_data.items()] + curve_data.sort(key=lambda x: x[0]) + steps = [step for step, value in curve_data] + value = [value for step, value in curve_data] + k_y = [{'data': np.array(value), 'label': label}] + save_path = os.path.join(save_folder, 'detail', f"{label}.png") + with FS.put_to(save_path) as local_file: + plot_multi_curves(x=np.array(steps), + y=k_y, + x_label='steps', + y_label=None, + title=f"{label}'s validation loss", + save_path=local_file) + y = y + k_y + if len(steps) > 0: + save_path = os.path.join(save_folder, f"summary.png") # noqa + with FS.put_to(save_path) as local_file: + plot_multi_curves( + x=np.array(steps), + y=y, + x_label='steps', + y_label=None, + title=f"validation loss", # noqa + save_path=local_file) + + def after_iter(self, solver): + if solver.mode == 'train' and solver.total_iter % self.val_interval == 0: + step = solver.total_iter + solver.eval_mode() + self.get_val_loss(solver, step) + solver.train_mode() + torch.cuda.synchronize() + barrier() + + def after_all_iter(self, solver): + if solver.mode == 'train': + step = solver.total_iter + solver.eval_mode() + self.get_val_loss(solver, step) + solver.train_mode() + torch.cuda.synchronize() + barrier() + + @staticmethod + def get_config_template(): + return dict_to_yaml('HOOK', + __class__.__name__, + ValLossHook.para_dict, + set_name=True) diff --git a/scepter/modules/solver/registry.py b/scepter/modules/solver/registry.py index 82919fd..dc0779d 100644 --- a/scepter/modules/solver/registry.py +++ b/scepter/modules/solver/registry.py @@ -17,8 +17,11 @@ def build_solver(cfg, registry, logger=None, *args, **kwargs): f'registry must be type Registry, got {type(registry)}') cfg = deep_copy(cfg) - req_type = cfg.get('NAME') + + from scepter.modules.utils.import_utils import LazyImportModule + sig = (registry.name.upper(), req_type) + LazyImportModule.import_module(sig) if isinstance(req_type, str): req_type_entry = registry.get(req_type) if req_type_entry is None: diff --git a/scepter/modules/transform/__init__.py b/scepter/modules/transform/__init__.py index eeb08c0..784f3a2 100644 --- a/scepter/modules/transform/__init__.py +++ b/scepter/modules/transform/__init__.py @@ -1,27 +1,59 @@ # -*- coding: utf-8 -*- # Copyright (c) Alibaba, Inc. and its affiliates. +from typing import TYPE_CHECKING +from scepter.modules.utils.import_utils import LazyImportModule -from scepter.modules.transform.augmention import ColorJitterGeneral -from scepter.modules.transform.compose import Compose -from scepter.modules.transform.identity import Identity -from scepter.modules.transform.image import (CenterCrop, FlexibleCenterCrop, - FlexibleResize, ImageToTensor, - ImageTransform, Normalize, - RandomHorizontalFlip, - RandomResizedCrop, Resize) -from scepter.modules.transform.io import (LoadCvImageFromFile, - LoadImageFromFile, - LoadImageFromFileList, - LoadPILImageFromFile) -from scepter.modules.transform.io_video import (DecodeVideoToTensor, - LoadVideoFromFile) -from scepter.modules.transform.registry import TRANSFORMS, build_pipeline -from scepter.modules.transform.tensor import (Rename, RenameMeta, Select, - TemplateStr, ToNumpy, ToTensor) -from scepter.modules.transform.transform_xl import FlexibleCropXL -from scepter.modules.transform.video import (AutoResizedCropVideo, - CenterCropVideo, NormalizeVideo, - RandomHorizontalFlipVideo, - RandomResizedCropVideo, - ResizeVideo, VideoToTensor, - VideoTransform) + +if TYPE_CHECKING: + from scepter.modules.transform.augmention import ColorJitterGeneral + from scepter.modules.transform.compose import Compose + from scepter.modules.transform.identity import Identity + from scepter.modules.transform.image import (CenterCrop, FlexibleCenterCrop, + FlexibleResize, ImageToTensor, + ImageTransform, Normalize, + RandomHorizontalFlip, + RandomResizedCrop, Resize) + from scepter.modules.transform.io import (LoadCvImageFromFile, + LoadImageFromFile, + LoadImageFromFileList, + LoadPILImageFromFile) + from scepter.modules.transform.io_video import (DecodeVideoToTensor, + LoadVideoFromFile) + from scepter.modules.transform.registry import TRANSFORMS, build_pipeline + from scepter.modules.transform.tensor import (Rename, RenameMeta, Select, + TemplateStr, ToNumpy, ToTensor) + from scepter.modules.transform.transform_xl import FlexibleCropXL + from scepter.modules.transform.video import (AutoResizedCropVideo, + CenterCropVideo, NormalizeVideo, + RandomHorizontalFlipVideo, + RandomResizedCropVideo, + ResizeVideo, VideoToTensor, + VideoTransform) +else: + _import_structure = { + 'augmention': ['ColorJitterGeneral'], + 'compose': ['Compose'], + 'identity': ['Identity'], + 'image': ['CenterCrop', 'FlexibleCenterCrop', 'FlexibleResize', + 'ImageToTensor', 'ImageTransform', 'Normalize', + 'RandomHorizontalFlip', 'RandomResizedCrop', 'Resize'], + 'io': ['LoadCvImageFromFile', 'LoadImageFromFile', + 'LoadImageFromFileList', 'LoadPILImageFromFile'], + 'io_video': ['DecodeVideoToTensor', 'LoadVideoFromFile'], + 'registry': ['TRANSFORMS', 'build_pipeline'], + 'tensor': ['Rename', 'RenameMeta', 'Select', 'TemplateStr', + 'ToNumpy', 'ToTensor'], + 'transform_xl': ['FlexibleCropXL'], + 'video': ['AutoResizedCropVideo', 'CenterCropVideo', 'NormalizeVideo', + 'RandomHorizontalFlipVideo', 'RandomResizedCropVideo', + 'ResizeVideo', 'VideoToTensor', 'VideoTransform'] + } + + import sys + sys.modules[__name__] = LazyImportModule( + __name__, + globals()['__file__'], + _import_structure, + module_spec=__spec__, + extra_objects={}, + ) diff --git a/scepter/modules/utils/__init__.py b/scepter/modules/utils/__init__.py index 8e97972..ace81a5 100644 --- a/scepter/modules/utils/__init__.py +++ b/scepter/modules/utils/__init__.py @@ -1,4 +1,23 @@ # -*- coding: utf-8 -*- # Copyright (c) Alibaba, Inc. and its affiliates. -from scepter.modules.utils import (config, distribute, file_clients, - file_system, module_transform) +from typing import TYPE_CHECKING +from scepter.modules.utils.import_utils import LazyImportModule + + +if TYPE_CHECKING: + from scepter.modules.utils import (config, distribute, file_clients, + file_system, module_transform) +else: + _import_structure = { + 'utils': ['config', 'distribute', 'file_clients', + 'file_system', 'module_transform'] + } + + import sys + sys.modules[__name__] = LazyImportModule( + __name__, + globals()['__file__'], + _import_structure, + module_spec=__spec__, + extra_objects={}, + ) diff --git a/scepter/modules/utils/ast_utils.py b/scepter/modules/utils/ast_utils.py new file mode 100644 index 0000000..9ab627c --- /dev/null +++ b/scepter/modules/utils/ast_utils.py @@ -0,0 +1,486 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +import ast +import logging +import os +import os.path as osp +import time +import traceback +from pathlib import Path +from typing import Union, Any + + +p = Path(__file__) + + +SKIP_FUNCTION_SCANNING = True +SCEPTER_PATH = p.resolve().parents[2] +REGISTER_CLASS = 'register_class' +IGNORED_PACKAGES = ['.'] +SCAN_SUB_FOLDERS = [ + 'modules', 'studio', 'tools', 'workflow' +] +INDEXER_FILE = 'ast_indexer' +DECORATOR_KEY = 'decorators' +EXPRESS_KEY = 'express' +FROM_IMPORT_KEY = 'from_imports' +IMPORT_KEY = 'imports' +FILE_NAME_KEY = 'filepath' +INDEX_KEY = 'index' +REQUIREMENT_KEY = 'requirements' +MODULE_KEY = 'module' +CLASS_NAME = 'class_name' + + +def get_ast_logger(): + ast_logger = logging.getLogger('scepter.ast') + ast_logger.setLevel(logging.INFO) + return ast_logger + + +logger = get_ast_logger() + + +class AstScanning(object): + def __init__(self) -> None: + self.result_import = dict() + self.result_from_import = dict() + self.result_decorator = [] + self.express = [] + + def _is_sub_node(self, node: object) -> bool: + return isinstance(node, + ast.AST) and not isinstance(node, ast.expr_context) + + def _is_leaf(self, node: ast.AST) -> bool: + for field in node._fields: + attr = getattr(node, field) + if self._is_sub_node(attr): + return False + elif isinstance(attr, (list, tuple)): + for val in attr: + if self._is_sub_node(val): + return False + else: + return True + + def _skip_function(self, node: Union[ast.AST, 'str']) -> bool: + if SKIP_FUNCTION_SCANNING: + if type(node).__name__ == 'FunctionDef' or node == 'FunctionDef': + return True + return False + + def _fields(self, n: ast.AST, show_offsets: bool = True) -> tuple: + if show_offsets: + return n._attributes + n._fields + else: + return n._fields + + def _leaf(self, node: ast.AST, show_offsets: bool = True) -> str: + output = dict() + if isinstance(node, ast.AST): + local_dict = dict() + for field in self._fields(node, show_offsets=show_offsets): + field_output = self._leaf( + getattr(node, field), show_offsets=show_offsets) + local_dict[field] = field_output + output[type(node).__name__] = local_dict + return output + else: + return node + + def _refresh(self): + self.result_import = dict() + self.result_from_import = dict() + self.result_decorator = [] + self.result_express = [] + + def scan_ast(self, node: Union[ast.AST, None, str]): + self._setup_global() + self.scan_import(node, indent=' ', show_offsets=False) + + def scan_import( + self, + node: Union[ast.AST, None, str], + show_offsets: bool = True, + parent_node_name: str = '', + ) -> None | str | dict[Any, Any]: + if node is None: + return node + elif self._is_leaf(node): + return self._leaf(node, show_offsets=show_offsets) + else: + def _scan_import(el: Union[ast.AST, None, str], + parent_node_name: str = '') -> str: + return self.scan_import( + el, + show_offsets=show_offsets, + parent_node_name=parent_node_name) + + outputs = dict() + # add relative path expression + if type(node).__name__ == 'ImportFrom': + level = getattr(node, 'level') + if level >= 1: + path_level = ''.join(['.'] * level) + setattr(node, 'level', 0) + module_name = getattr(node, 'module') + if module_name is None: + setattr(node, 'module', path_level) + else: + setattr(node, 'module', path_level + module_name) + + for field in self._fields(node, show_offsets=show_offsets): + attr = getattr(node, field) + if not attr: + outputs[field] = [] + elif self._skip_function(parent_node_name): + continue + elif (isinstance(attr, list) and len(attr) == 1 + and isinstance(attr[0], ast.AST) + and self._is_leaf(attr[0])): + local_out = _scan_import(attr[0]) + outputs[field] = local_out + elif isinstance(attr, list): + el_dict = dict() + for el in attr: + local_out = _scan_import(el, type(el).__name__) + name = type(el).__name__ + if (name == 'Import' or name == 'ImportFrom' + or parent_node_name == 'ImportFrom' + or parent_node_name == 'Import'): + if name not in el_dict: + el_dict[name] = [] + el_dict[name].append(local_out) + outputs[field] = el_dict + elif isinstance(attr, ast.AST): + output = _scan_import(attr) + outputs[field] = output + else: + outputs[field] = attr + + if (type(node).__name__ == 'Import' + or type(node).__name__ == 'ImportFrom'): + if type(node).__name__ == 'ImportFrom': + if field == 'module': + self.result_from_import[outputs[field]] = dict() + if field == 'names': + if isinstance(outputs[field]['alias'], list): + item_name = [] + for item in outputs[field]['alias']: + local_name = item['alias']['name'] + item_name.append(local_name) + self.result_from_import[ + outputs['module']] = item_name + else: + local_name = outputs[field]['alias']['name'] + self.result_from_import[outputs['module']] = [ + local_name + ] + + if type(node).__name__ == 'Import': + final_dict = outputs[field]['alias'] + if isinstance(final_dict, list): + for item in final_dict: + self.result_import[item['alias'] + ['name']] = item['alias'] + else: + self.result_import[outputs[field]['alias'] + ['name']] = final_dict + + if 'decorator_list' == field and attr != []: + for item in attr: + setattr(item, CLASS_NAME, node.name) + self.result_decorator.extend(attr) + + if attr != [] and type( + attr + ).__name__ == 'Call' and parent_node_name == 'Expr': + self.result_express.append(attr) + return {IMPORT_KEY: self.result_import, + FROM_IMPORT_KEY: self.result_from_import, + DECORATOR_KEY: self.result_decorator, + EXPRESS_KEY: self.result_express} + + def _parse_decorator(self, node: ast.AST) -> tuple: + def _get_attribute_item(node: ast.AST) -> tuple: + value, id, attr = None, None, None + if type(node).__name__ == 'Attribute': + value = getattr(node, 'value') + id = getattr(value, 'id', None) + attr = getattr(node, 'attr') + if type(node).__name__ == 'Name': + id = getattr(node, 'id') + return id, attr + + def _get_args_name(nodes: list) -> list: + result = [] + for node in nodes: + if type(node).__name__ == 'Str': + result.append((node.s, None)) + elif type(node).__name__ == 'Constant': + result.append((node.value, None)) + else: + result.append(_get_attribute_item(node)) + return result + + def _get_keyword_name(nodes: ast.AST) -> list: + result = [] + for node in nodes: + if type(node).__name__ == 'keyword': + attribute_node = getattr(node, 'value') + if type(attribute_node).__name__ == 'Str': + result.append((getattr(node, + 'arg'), attribute_node.s, None)) + elif type(attribute_node).__name__ == 'Constant': + result.append( + (getattr(node, 'arg'), attribute_node.value, None)) + else: + result.append((getattr(node, 'arg'), ) + + _get_attribute_item(attribute_node)) + return result + + functions = _get_attribute_item(node.func) + args_list = _get_args_name(node.args) + keyword_list = _get_keyword_name(node.keywords) + return functions, args_list, keyword_list + + def _registry_indexer(self, parsed_input: tuple, class_name: str) -> tuple: + """format registry information to a tuple indexer + + Return: + tuple: (MODELS, ClassName, RegisterName) + """ + functions, args_list, keyword_list = parsed_input + + if REGISTER_CLASS != functions[1]: + return None + output = [functions[0]] + return (output[0], class_name, args_list) + + def parse_decorators(self, nodes: list) -> list: + """parse the AST nodes of decorators object to registry indexer + + Args: + nodes (list): list of AST decorator nodes + + Returns: + list: list of registry indexer + """ + results = [] + for node in nodes: + if type(node).__name__ != 'Call': + continue + class_name = getattr(node, CLASS_NAME, None) + func = getattr(node, 'func') + if getattr(func, 'attr', None) != REGISTER_CLASS: + continue + + parse_output = self._parse_decorator(node) + index = self._registry_indexer(parse_output, class_name) + if None is not index: + results.append(index) + return results + + def generate_ast(self, file): + self._refresh() + with open(file, 'r', encoding='utf8') as code: + data = code.readlines() + data = ''.join(data) + node = ast.parse(data) + output = self.scan_import(node, show_offsets=False) + output[DECORATOR_KEY] = self.parse_decorators(output[DECORATOR_KEY]) + output[EXPRESS_KEY] = self.parse_decorators(output[EXPRESS_KEY]) + output[DECORATOR_KEY].extend(output[EXPRESS_KEY]) + return output + + +class FilesAstScanning(object): + def __init__(self) -> None: + self.astScaner = AstScanning() + self.file_dirs = [] + self.requirement_dirs = [] + + def _parse_import_path(self, + import_package: str, + current_path: str = None) -> str: + """ + Args: + import_package (str): relative import or abs import + current_path (str): path/to/current/file + """ + if import_package.startswith(IGNORED_PACKAGES[0]): + return SCEPTER_PATH + '/' + '/'.join( + import_package.split('.')[1:]) + '.py' + elif import_package.startswith(IGNORED_PACKAGES[1]): + current_path_list = current_path.split('/') + import_package_list = import_package.split('.') + level = 0 + for index, item in enumerate(import_package_list): + if item != '': + level = index + break + + abs_path_list = current_path_list[0:-level] + abs_path_list.extend(import_package_list[index:]) + return '/' + '/'.join(abs_path_list) + '.py' + else: + return current_path + + def parse_import(self, scan_result: dict) -> list: + """parse import and from import dicts to a third party package list + + Args: + scan_result (dict): including the import and from import result + + Returns: + list: a list of package ignored 'scepter' and relative path import + """ + output = [] + output.extend(list(scan_result[IMPORT_KEY].keys())) + output.extend(list(scan_result[FROM_IMPORT_KEY].keys())) + + # get the package name + for index, item in enumerate(output): + if '' == item.split('.')[0]: + output[index] = '.' + else: + output[index] = item.split('.')[0] + + ignored = set() + for item in output: + for ignored_package in IGNORED_PACKAGES: + if item.startswith(ignored_package): + ignored.add(item) + return list(set(output) - set(ignored)) + + def traversal_files(self, path, check_sub_dir=None, include_init=False): + self.file_dirs = [] + if check_sub_dir is None or len(check_sub_dir) == 0: + self._traversal_files(path, include_init=include_init) + else: + for item in check_sub_dir: + sub_dir = os.path.join(path, item) + if os.path.isdir(sub_dir): + self._traversal_files(sub_dir, include_init=include_init) + + def _traversal_files(self, path, include_init=False): + dir_list = os.scandir(path) + for item in dir_list: + if item.name == '__init__.py' and not include_init: + continue + elif (item.name.startswith('__') + and item.name != '__init__.py') or item.name.endswith( + '.json') or item.name.endswith('.md'): + continue + if item.is_dir(): + self._traversal_files(item.path, include_init=include_init) + elif item.is_file() and item.name.endswith('.py'): + self.file_dirs.append(item.path) + elif item.is_file() and 'requirement' in item.name: + self.requirement_dirs.append(item.path) + + def _get_single_file_scan_result(self, file): + try: + output = self.astScaner.generate_ast(file) + except Exception as e: + detail = traceback.extract_tb(e.__traceback__) + raise Exception( + f'During ast indexing the file {file}, a related error excepted ' + f'in the file {detail[-1].filename} at line: ' + f'{detail[-1].lineno}: "{detail[-1].line}" with error msg: ' + f'"{type(e).__name__}: {e}", please double check the origin file {file} ' + f'to see whether the file is correctly edited.') + + import_list = self.parse_import(output) + return output[DECORATOR_KEY], import_list + + def _inverted_index(self, forward_index): + inverted_index = dict() + for index in forward_index: + for item in forward_index[index][DECORATOR_KEY]: + inverted_index[item[:2]] = { + FILE_NAME_KEY: index, + IMPORT_KEY: forward_index[index][IMPORT_KEY], + MODULE_KEY: forward_index[index][MODULE_KEY], + } + if item[-1]: + for register_name in item[-1]: + inverted_index[(item[0], register_name[0])] = { + FILE_NAME_KEY: index, + IMPORT_KEY: forward_index[index][IMPORT_KEY], + MODULE_KEY: forward_index[index][MODULE_KEY], + } + return inverted_index + + def _module_import(self, forward_index): + module_import = dict() + for index, value_dict in forward_index.items(): + module_import[value_dict[MODULE_KEY]] = value_dict[IMPORT_KEY] + return module_import + + def get_files_scan_results(self, + target_file_list=None, + target_dir=SCEPTER_PATH, + target_folders=SCAN_SUB_FOLDERS): + """the entry method of the ast scan method + + Args: + target_file_list can override the dir and folders combine + target_dir (str, optional): the absolute path of the target directory to be scanned. Defaults to None. + target_folder (list, optional): the list of + sub-folders to be scanned in the target folder. + Defaults to SCAN_SUB_FOLDERS. + + Returns: + dict: indexer of registry + """ + start = time.time() + if target_file_list is not None: + self.file_dirs = target_file_list + else: + self.traversal_files(target_dir, target_folders) + logger.info( + f'AST-Scanning the path "{target_dir}" with the following sub folders {target_folders}' + ) + + result = dict() + for file in self.file_dirs: + filepath = file[file.rfind('scepter'):] + module_name = filepath.replace(osp.sep, '.').replace('.py', '') + decorator_list, import_list = self._get_single_file_scan_result( + file) + result[file] = { + DECORATOR_KEY: decorator_list, + IMPORT_KEY: import_list, + MODULE_KEY: module_name + } + + inverted_index_with_results = self._inverted_index(result) + module_import = self._module_import(result) + index = { + INDEX_KEY: inverted_index_with_results, + REQUIREMENT_KEY: module_import + } + logger.info( + f'Scanning done! A number of {len(inverted_index_with_results)} ' + f'components indexed or updated! Time consumed {time.time()-start}s' + ) + return index + + +file_scanner = FilesAstScanning() +file_index = None + + +def load_index(file_list=None): + global file_index + if file_index is None: + logger.info('Building ast index from scanning every file!') + file_index = file_scanner.get_files_scan_results(file_list) + return file_index + + +if __name__ == '__main__': + index = load_index() + print(index) diff --git a/scepter/modules/utils/config.py b/scepter/modules/utils/config.py index 193292a..15cdd65 100644 --- a/scepter/modules/utils/config.py +++ b/scepter/modules/utils/config.py @@ -10,7 +10,8 @@ import sys import yaml -from scepter.modules.utils.model import StdMsg +from scepter.modules.utils.logger import StdMsg + _SECURE_KEYWORDS = [ 'ENDPOINT', 'BUCKET', 'OSS_AK', 'OSS_SK', 'OSS', 'TOKEN', 'APPKEY' @@ -211,7 +212,7 @@ def dict_to_yaml(module_name, name, json_config, set_name=False, exclude_keys=[] return yaml_str -pattern = re.compile('.*?(\${\w+}).*?') # noqa +pattern = re.compile(r'.*?(\${\w+}).*?') # noqa def env_var_constructor(loader, node): @@ -669,4 +670,4 @@ class Config(object): return len(self.cfg_dict) def pop(self, name): - self.cfg_dict.pop(name) + return self.cfg_dict.pop(name) diff --git a/scepter/modules/utils/distribute.py b/scepter/modules/utils/distribute.py index 53d5a14..e54465f 100644 --- a/scepter/modules/utils/distribute.py +++ b/scepter/modules/utils/distribute.py @@ -13,8 +13,8 @@ import numpy as np import torch import torch.distributed as dist from torch.autograd import Function - -from scepter.modules.utils.model import StdMsg +import platform +from scepter.modules.utils.logger import StdMsg __all__ = [ 'gather_data', 'we', 'broadcast', 'barrier', 'reduce_scatter', 'reduce', @@ -620,7 +620,11 @@ class Workenv(object): self.sync_bn = False self.rank = 0 self.world_size = 1 - self.device_id = 0 + if torch.cuda.is_available(): + self.device_id = 0 + else: + self.device_id = 'mps' if platform.system() == "Darwin" else 'cpu' + self.backend = '' self.device_count = 1 self.seed = 2023 @@ -665,7 +669,7 @@ class Workenv(object): self.share_storage = os.environ.get('SHARE_STORAGE', None) == 'true' if not torch.cuda.is_available(): - self.device_id = 'cpu' + self.device_id = 'mps' if platform.system() == "Darwin" else 'cpu' fn(config) return diff --git a/scepter/modules/utils/error.py b/scepter/modules/utils/error.py new file mode 100644 index 0000000..7b504c0 --- /dev/null +++ b/scepter/modules/utils/error.py @@ -0,0 +1,208 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. + +# docstyle-ignore +ALBUMENTATIONS_IMPORT_ERROR = """ +{0} requires the albumentations library but it was not found in your environment. You can install it with pip: +`pip install albumentations` +""" + +# docstyle-ignore +SENTENCEPIECE_IMPORT_ERROR = """ +{0} requires the SentencePiece library but it was not found in your environment. Checkout the instructions on the +installation page of its repo: https://github.com/google/sentencepiece#installation and follow the ones +that match your environment. +""" + +# docstyle-ignore +SKLEARN_IMPORT_ERROR = """ +{0} requires the scikit-learn library but it was not found in your environment. You can install it with: +``` +pip install -U scikit-learn +``` +In a notebook or a colab, you can install it by executing a cell with +``` +!pip install -U scikit-learn +``` +""" + +# docstyle-ignore +TIMM_IMPORT_ERROR = """ +{0} requires the timm library but it was not found in your environment. You can install it with pip: +`pip install timm` +""" + +# docstyle-ignore +SCEPTER_IMPORT_ERROR = """ +{0} requires the scepter library but it was not found in your environment. You can install it with pip: +`pip install scepter` +""" + +# docstyle-ignore +PYTORCH_IMPORT_ERROR = """ +{0} requires the PyTorch library but it was not found in your environment. Checkout the instructions on the +installation page: https://pytorch.org/get-started/locally/ and follow the ones that match your environment. +""" + +WENETRUNTIME_IMPORT_ERROR = """ +{0} requires the wenetruntime library but it was not found in your environment. You can install it with pip: +`pip install wenetruntime==TORCH_VER` +""" + +# docstyle-ignore +TORCHVISION_IMPORT_ERROR = """ +{0} requires the scipy library but it was not found in your environment. You can install it with pip: +`pip install torchvision` +""" + +# docstyle-ignore +OPENCV_IMPORT_ERROR = """ +{0} requires the opencv library but it was not found in your environment. You can install it with pip: +`pip install opencv-python` +""" + +PILLOW_IMPORT_ERROR = """ +{0} requires the Pillow library but it was not found in your environment. You can install it with pip: +`pip install Pillow` +""" + +MODELSCOPE_IMPORT_ERROR = """ +{0} requires the modelscope library but it was not found in your environment. You can install it with pip: +`pip install modelscope` +""" + +FLASH_ATTN_IMPORT_ERROR = """ +{0} requires the flash_attn library but it was not found in your environment. You can install it with pip: +`pip install flash_attn==2.5.8` +""" + +XFORMERS_IMPORT_ERROR = """ +{0} requires the xformers library but it was not found in your environment. You can install it with pip: +`pip install xformers` +""" + +DECORD_IMPORT_ERROR = """ +{0} requires the decord library but it was not found in your environment. You can install it with pip: +`pip install decord>=0.6.0` +""" + +# docstyle-ignore +BEAUTIFULSOUP4_IMPORT_ERROR = """ +{0} requires the decord library but it was not found in your environment. You can install it with pip: +`pip install beautifulsoup4` +""" + +# docstyle-ignore +BEZIER_IMPORT_ERROR = """ +{0} requires the beizer library but it was not found in your environment. You can install it with pip: +`pip install beizer` +""" + +# docstyle-ignore +EINOPS_IMPORT_ERROR = """ +{0} requires the einops library but it was not found in your environment. You can install it with pip: +`pip install einops` +""" + +# docstyle-ignore +EASYNLP_IMPORT_ERROR = """ +{0} requires the easynlp library but it was not found in your environment. +You can install it with pip on linux or mac: +`pip install pai-easynlp -f https://modelscope.oss-cn-beijing.aliyuncs.com/releases/repo.html` +Or you can checkout the instructions on the +installation page: https://github.com/alibaba/EasyNLP and follow the ones that match your environment. +""" + +# docstyle-ignore +NUMPY_IMPORT_ERROR = """ +{0} requires the megatron_util library but it was not found in your environment. You can install it with pip: +`pip install numpy` +""" + +# docstyle-ignore +OSS2_IMPORT_ERROR = """ +{0} requires the oss2 library but it was not found in your environment. You can install it with pip: +`pip install oss2` +""" + +# docstyle-ignore +PYCOCOTOOLS_IMPORT_ERROR = """ +{0} requires the pycocotools library but it was not found in your environment. You can install it with pip: +`pip install pycocotools` +""" + +# docstyle-ignore +OPENCLIP_IMPORT_ERROR = """ +{0} requires the fasttext library but it was not found in your environment. +You can install it with pip on linux or mac: +`pip install open_clip_torch` +Or you can checkout the instructions on the +installation page: https://github.com/mlfoundations/open_clip and follow the ones that match your environment. +""" + +# docstyle-ignore +PYYAML_IMPORT_ERROR = """ +{0} requires the pyyaml library but it was not found in your environment. You can install it with pip: +`pip install pyyaml` +""" + +# docstyle-ignore +SWIFT_IMPORT_ERROR = """ +{0} requires the ms-swift library but it was not found in your environment. You can install it with pip: +`pip install ms-swift` +""" + +SCIKIT_IMAGE_IMPORT_ERROR = """ +{0} requires the ms-swift library but it was not found in your environment. You can install it with pip: +`pip install scikit-image` +""" + +SCIKIT_LEARN_IMPORT_ERROR = """ +{0} requires the ms-swift library but it was not found in your environment. You can install it with pip: +`pip install scikit-learn` +""" + +TORCHSDE_IMPORT_ERROR = """ +{0} requires the ms-swift library but it was not found in your environment. You can install it with pip: +`pip install torchsde +""" + +BITSANDBYTES_IMPORT_ERROR = """ +{0} requires the ms-swift library but it was not found in your environment. You can install it with pip: +`pip install bitsandbytes +""" + +GRADIO_IMAGESLIDER_IMPORT_ERROR = """ +{0} requires the ms-swift library but it was not found in your environment. You can install it with pip: +`pip install gradio_imageslider +""" + +IMAGEHASH_IMPORT_ERROR = """ +{0} requires the ms-swift library but it was not found in your environment. You can install it with pip: +`pip install imagehash +""" + +PSUTIL_IMPORT_ERROR = """ +{0} requires the ms-swift library but it was not found in your environment. You can install it with pip: +`pip install psutil +""" + +TIKTOKEN_IMPORT_ERROR = """ +{0} requires the ms-swift library but it was not found in your environment. You can install it with pip: +`pip install tiktoken +""" + +TRANSFORMERS_IMPORT_ERROR = """ +{0} requires the transformers library but it was not found in your environment. You can install it with pip: +`pip install transformers` +""" + +GENERAL_IMPORT_ERROR = """ +{0} requires the REQ library but it was not found in your environment. You can install it with pip: +`pip install REQ` +""" + +GRADIO_IMPORT_ERROR = """ +{0} requires the gradio library but it was not found in your environment. You can install it with pip: +`pip install gradio` +""" \ No newline at end of file diff --git a/scepter/modules/utils/file_clients/__init__.py b/scepter/modules/utils/file_clients/__init__.py index 52fcbd5..71fcb03 100644 --- a/scepter/modules/utils/file_clients/__init__.py +++ b/scepter/modules/utils/file_clients/__init__.py @@ -1,7 +1,29 @@ # -*- coding: utf-8 -*- # Copyright (c) Alibaba, Inc. and its affiliates. -from scepter.modules.utils.file_clients.aliyun_oss_fs import AliyunOssFs -from scepter.modules.utils.file_clients.http_fs import HttpFs -from scepter.modules.utils.file_clients.huggingface_fs import HuggingfaceFs -from scepter.modules.utils.file_clients.local_fs import LocalFs -from scepter.modules.utils.file_clients.modelscope_fs import ModelscopeFs +from typing import TYPE_CHECKING +from scepter.modules.utils.import_utils import LazyImportModule + + +if TYPE_CHECKING: + from scepter.modules.utils.file_clients.aliyun_oss_fs import AliyunOssFs + from scepter.modules.utils.file_clients.http_fs import HttpFs + from scepter.modules.utils.file_clients.huggingface_fs import HuggingfaceFs + from scepter.modules.utils.file_clients.local_fs import LocalFs + from scepter.modules.utils.file_clients.modelscope_fs import ModelscopeFs +else: + _import_structure = { + 'aliyun_oss_fs': ['AliyunOssFs'], + 'http_fs': ['HttpFs'], + 'huggingface_fs': ['HuggingfaceFs'], + 'local_fs': ['LocalFs'], + 'modelscope_fs': ['ModelscopeFs'] + } + + import sys + sys.modules[__name__] = LazyImportModule( + __name__, + globals()['__file__'], + _import_structure, + module_spec=__spec__, + extra_objects={}, + ) diff --git a/scepter/modules/utils/import_utils.py b/scepter/modules/utils/import_utils.py new file mode 100644 index 0000000..e66c6e2 --- /dev/null +++ b/scepter/modules/utils/import_utils.py @@ -0,0 +1,333 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +import functools +import importlib +import logging +import os +import sys +from collections import OrderedDict +from importlib import import_module +from itertools import chain +from types import ModuleType +from typing import Any +import scepter +from scepter.modules.utils.ast_utils import (INDEX_KEY, + MODULE_KEY, + REQUIREMENT_KEY, + load_index) +from scepter.modules.utils.error import * +from scepter.modules.utils.logger import get_logger + + +if sys.version_info < (3, 8): + import importlib_metadata +else: + import importlib.metadata as importlib_metadata + +logger = get_logger() + + +def get_dirname(): + return os.path.dirname(scepter.__file__) + + +def import_modules(imports, allow_failed_imports=False): + """Import modules from the given list of strings. + + Args: + imports (list | str | None): The given module names to be imported. + allow_failed_imports (bool): If True, the failed imports will return + None. Otherwise, an ImportError is raise. Default: False. + + Returns: + list[module] | module | None: The imported modules. + + Examples: + >>> osp, sys = import_modules( + ... ['os.path', 'sys']) + >>> import os.path as osp_ + >>> import sys as sys_ + >>> assert osp == osp_ + >>> assert sys == sys_ + """ + if not imports: + return + single_import = False + if isinstance(imports, str): + single_import = True + imports = [imports] + if not isinstance(imports, list): + raise TypeError( + f'custom_imports must be a list but got type {type(imports)}') + imported = [] + for imp in imports: + if not isinstance(imp, str): + raise TypeError( + f'{imp} is of type {type(imp)} and cannot be imported.') + try: + imported_tmp = import_module(imp) + except ImportError: + if allow_failed_imports: + logger.warning(f'{imp} failed to import and is ignored.') + imported_tmp = None + else: + raise ImportError + imported.append(imported_tmp) + if single_import: + imported = imported[0] + return imported + + +# following code borrows implementation from huggingface/transformers +ENV_VARS_TRUE_VALUES = {'1', 'ON', 'YES', 'TRUE'} +ENV_VARS_TRUE_AND_AUTO_VALUES = ENV_VARS_TRUE_VALUES.union({'AUTO'}) +USE_TORCH = os.environ.get('USE_TORCH', 'AUTO').upper() + +_torch_version = 'N/A' +if USE_TORCH in ENV_VARS_TRUE_AND_AUTO_VALUES: + _torch_available = importlib.util.find_spec('torch') is not None + if _torch_available: + try: + _torch_version = importlib_metadata.version('torch') + logger.info(f'PyTorch version {_torch_version} Found.') + except importlib_metadata.PackageNotFoundError: + _torch_available = False +else: + logger.info('Disabling PyTorch because USE_TF is set') + _torch_available = False + + +def is_torchvision_available(): + return importlib.util.find_spec('torchvision') is not None + + +def is_sentencepiece_available(): + return importlib.util.find_spec('sentencepiece') is not None + + +def is_scepter_available(): + return importlib.util.find_spec('scepter') is not None + + +def is_torch_available(): + return _torch_available + + +def is_torch_cuda_available(): + if is_torch_available(): + import torch + return torch.cuda.is_available() + else: + return False + + +def is_swift_available(): + return importlib.util.find_spec('swift') is not None + + +def is_opencv_available(): + return importlib.util.find_spec('cv2') is not None + + +def is_pillow_available(): + return importlib.util.find_spec('PIL.Image') is not None + + +def _is_package_available_fn(pkg_name): + return importlib.util.find_spec(pkg_name) is not None + + +def is_package_available(pkg_name): + return functools.partial(_is_package_available_fn, pkg_name) + + +def is_flash_attn_available(): + return importlib.util.find_spec('flash-attn') is not None + + +def is_transformers_available(): + return importlib.util.find_spec('transformers') is not None + + +REQUIREMENTS_MAAPING = OrderedDict([ + ('scepter', (is_scepter_available(), SCEPTER_IMPORT_ERROR)), + ('torch', (is_torch_available, PYTORCH_IMPORT_ERROR)), + ('torchvision', (is_torchvision_available(), TORCHVISION_IMPORT_ERROR)), + ('cv2', (is_opencv_available, OPENCV_IMPORT_ERROR)), + ('PIL', (is_pillow_available, PILLOW_IMPORT_ERROR)), + ('modelscope', (is_package_available('modelscope'), MODELSCOPE_IMPORT_ERROR)), + ('flash-attn', (is_flash_attn_available, FLASH_ATTN_IMPORT_ERROR)), + ('xformers', (is_package_available('funasr'), XFORMERS_IMPORT_ERROR)), + ('albumentations', (is_package_available('albumentations'), ALBUMENTATIONS_IMPORT_ERROR)), + ('decord', (is_package_available('decord'), DECORD_IMPORT_ERROR)), + ('beautifulsoup4', (is_package_available('beautifulsoup4'), BEAUTIFULSOUP4_IMPORT_ERROR)), + ('bezier', (is_package_available('bezier'), BEZIER_IMPORT_ERROR)), + ('einops', (is_package_available('einops'), EINOPS_IMPORT_ERROR)), + ('numpy', (is_package_available('numpy'), NUMPY_IMPORT_ERROR)), + ('oss2', (is_package_available('oss2'), OSS2_IMPORT_ERROR)), + ('pycocotools', (is_package_available('pycocotools'), PYCOCOTOOLS_IMPORT_ERROR)), + ('open_clip', (is_package_available('open_clip'), OPENCLIP_IMPORT_ERROR)), + ('pyyaml', (is_package_available('pyyaml'), PYYAML_IMPORT_ERROR)), + ('transformers', (is_package_available('transformers'), TRANSFORMERS_IMPORT_ERROR)), + ('ms-swift', (is_package_available('ms-swift'), SWIFT_IMPORT_ERROR)), + ('gradio', (is_package_available('gradio'), SWIFT_IMPORT_ERROR)), + ('scikit-image', (is_package_available('scikit-image'), SCIKIT_IMAGE_IMPORT_ERROR)), + ('scikit-learn', (is_package_available('scikit-learn'), SCIKIT_LEARN_IMPORT_ERROR)), + ('sentencepiece', (is_package_available('sentencepiece'), SENTENCEPIECE_IMPORT_ERROR)), + ('torchsde', (is_package_available('torchsde'), TORCHSDE_IMPORT_ERROR)), + ('bitsandbytes', (is_package_available('bitsandbytes'), BITSANDBYTES_IMPORT_ERROR)), + ('gradio_imageslider', (is_package_available('gradio_imageslider'), GRADIO_IMAGESLIDER_IMPORT_ERROR)), + ('imagehash', (is_package_available('imagehash'), IMAGEHASH_IMPORT_ERROR)), + ('psutil', (is_package_available('psutil'), PSUTIL_IMPORT_ERROR)), + ('tiktoken', (is_package_available('tiktoken'), TIKTOKEN_IMPORT_ERROR)) +]) + +SYSTEM_PACKAGE = set(['os', 'sys', 'typing']) + + +def requires(obj, requirements): + if not isinstance(requirements, (list, tuple)): + requirements = [requirements] + if isinstance(obj, str): + name = obj + else: + name = obj.__name__ if hasattr(obj, + '__name__') else obj.__class__.__name__ + + checks = [] + for req in requirements: + if req == '' or req in SYSTEM_PACKAGE: + continue + if req in REQUIREMENTS_MAAPING: + check = REQUIREMENTS_MAAPING[req] + else: + check_fn = is_package_available(req) + err_msg = GENERAL_IMPORT_ERROR.replace('REQ', req) + check = (check_fn, err_msg) + checks.append(check) + + failed = [msg.format(name) for available, msg in checks if not available] + if failed: + raise ImportError(''.join(failed)) + + +def torch_required(func): + # Chose a different decorator name than in tests so it's clear they are not the same. + @functools.wraps(func) + def wrapper(*args, **kwargs): + if is_torch_available(): + return func(*args, **kwargs) + else: + raise ImportError(f'Method `{func.__name__}` requires PyTorch.') + return wrapper + + +class LazyImportModule(ModuleType): + _AST_INDEX = None + + def __init__(self, + name, + module_file, + import_structure, + module_spec=None, + extra_objects=None, + try_to_pre_import=False): + super().__init__(name) + self._modules = set(import_structure.keys()) + self._class_to_module = {} + for key, values in import_structure.items(): + for value in values: + self._class_to_module[value] = key + # Needed for autocompletion in an IDE + self.__all__ = list(import_structure.keys()) + list( + chain(*import_structure.values())) + self.__file__ = module_file + self.__spec__ = module_spec + self.__path__ = [os.path.dirname(module_file)] + self._objects = {} if extra_objects is None else extra_objects + self._name = name + self._import_structure = import_structure + if try_to_pre_import: + self._try_to_import() + + def _try_to_import(self): + for sub_module in self._class_to_module.keys(): + try: + getattr(self, sub_module) + except Exception as e: + logger.warning( + f'pre load module {sub_module} error, please check {e}') + + def __dir__(self): + result = super().__dir__() + for attr in self.__all__: + if attr not in result: + result.append(attr) + return result + + def __getattr__(self, name: str) -> Any: + if name in self._objects: + return self._objects[name] + if name in self._modules: + value = self._get_module(name) + elif name in self._class_to_module.keys(): + module = self._get_module(self._class_to_module[name]) + value = getattr(module, name) + else: + raise AttributeError( + f'module {self.__name__} has no attribute {name}') + + setattr(self, name, value) + return value + + def _get_module(self, module_name: str): + try: + module_name_full = self.__name__ + '.' + module_name + if not any( + module_name_full.startswith(f'scepter.{prefix}') + for prefix in ['modules', 'studio', 'version', 'tools', 'workflow']): + # check requirements before module import + requirements = self.get_requirements() + if module_name_full in requirements: + requires(module_name_full, requirements) + return importlib.import_module('.' + module_name, self.__name__) + except Exception as e: + raise RuntimeError( + f'Failed to import {self.__name__}.{module_name} because of the following error ' + f'(look up to see its traceback):\n{e}') from e + + def __reduce__(self): + return self.__class__, (self._name, self.__file__, + self._import_structure) + + @staticmethod + def get_ast_index(): + if LazyImportModule._AST_INDEX is None: + LazyImportModule._AST_INDEX = load_index() + return LazyImportModule._AST_INDEX + + @staticmethod + def import_module(signature): + """ import a lazy import module using signature + + Args: + signature (tuple): a tuple of str, (registry_name, class_name) + """ + ast_index = LazyImportModule.get_ast_index() + if signature in ast_index[INDEX_KEY]: + mod_index = ast_index[INDEX_KEY][signature] + module_name = mod_index[MODULE_KEY] + if module_name in ast_index[REQUIREMENT_KEY]: + requirements = ast_index[REQUIREMENT_KEY][module_name] + requires(module_name, requirements) + importlib.import_module(module_name) + else: + logger.warning(f'{signature} not found in ast index file') + + @staticmethod + def get_module_type(module): + ast_index = LazyImportModule.get_ast_index() + if module in ast_index[INDEX_KEY]: + return True + else: + return False diff --git a/scepter/modules/utils/logger.py b/scepter/modules/utils/logger.py index 99c4e97..50de436 100644 --- a/scepter/modules/utils/logger.py +++ b/scepter/modules/utils/logger.py @@ -1,15 +1,13 @@ # -*- coding: utf-8 -*- # Copyright (c) Alibaba, Inc. and its affiliates. - import logging -import numbers import sys import time from collections import OrderedDict import numpy as np import torch -from scepter.modules.utils.distribute import get_dist_info +import numbers def as_time(s): @@ -71,6 +69,7 @@ def init_logger(in_logger, log_file=None, dist_launcher='pytorch'): log_file (str, None): if not None, a file handler will be add to in_logger dist_launcher (str, None): """ + from scepter.modules.utils.distribute import get_dist_info rank, _ = get_dist_info() if rank == 0: if log_file is not None: @@ -93,6 +92,20 @@ def init_logger(in_logger, log_file=None, dist_launcher='pytorch'): in_logger.setLevel(logging.INFO) +class StdMsg(): + def __init__(self, name='msg'): + self.name = name + + def info(self, msg): + sys.stdout.write('[Info]: ' + msg + '\n') + + def error(self, msg): + sys.stdout.write('[Error]: ' + msg + '\n') + + def warning(self, msg): + sys.stdout.write('[Warning]: ' + msg + '\n') + + class LogAgg(object): """ Log variable aggregate tool. Recommend to invoke clear() function after one epoch. In distributed training environment, tensor variable will be all reduced to get an average. diff --git a/scepter/modules/utils/model.py b/scepter/modules/utils/model.py index 9ba924f..a8b96d5 100644 --- a/scepter/modules/utils/model.py +++ b/scepter/modules/utils/model.py @@ -2,7 +2,6 @@ # Copyright (c) Alibaba, Inc. and its affiliates. import os import re -import sys from collections import OrderedDict import torch @@ -10,20 +9,6 @@ import torch.nn as nn from torch.utils.model_zoo import load_url as load_state_dict_from_url -class StdMsg(): - def __init__(self, name='msg'): - self.name = name - - def info(self, msg): - sys.stdout.write('[Info]: ' + msg + '\n') - - def error(self, msg): - sys.stdout.write('[Error]: ' + msg + '\n') - - def warning(self, msg): - sys.stdout.write('[Warning]: ' + msg + '\n') - - def move_model_to_cpu(params): cpu_params = OrderedDict() for key, val in params.items(): @@ -41,7 +26,7 @@ def load_pretrained(model: torch.nn.Module, f'Load pretrained model [{model.__class__.__name__}] from {path}') if os.path.exists(path): # From local - state_dict = torch.load(path, map_location) + state_dict = torch.load(path, map_location, weights_only=True) elif path.startswith('http'): # From url state_dict = load_state_dict_from_url(path, diff --git a/scepter/modules/utils/registry.py b/scepter/modules/utils/registry.py index 39cae95..21eab6a 100644 --- a/scepter/modules/utils/registry.py +++ b/scepter/modules/utils/registry.py @@ -65,6 +65,13 @@ def build_from_config(cfg, registry, logger=None, *args, **kwargs): cfg = deep_copy(cfg) req_type = cfg.get('NAME') + + from scepter.modules.utils.import_utils import LazyImportModule + sig = (registry.name.upper(), req_type) + if (LazyImportModule.get_module_type(sig) + and req_type not in registry.class_map.keys()): + LazyImportModule.import_module(sig) + if isinstance(req_type, str): req_type_entry = registry.get(req_type) if req_type_entry is None: diff --git a/scepter/modules/utils/video_reader/__init__.py b/scepter/modules/utils/video_reader/__init__.py index 5392d67..3f71949 100644 --- a/scepter/modules/utils/video_reader/__init__.py +++ b/scepter/modules/utils/video_reader/__init__.py @@ -1,6 +1,26 @@ # -*- coding: utf-8 -*- # Copyright (c) Alibaba, Inc. and its affiliates. -from .frame_sampler import (FRAME_SAMPLERS, IntervalSampler, SegmentSampler, - UniformSampler, do_frame_sample) -from .video_reader import (EasyVideoReader, FramesReaderWrapper, - VideoReaderWrapper) +from typing import TYPE_CHECKING +from scepter.modules.utils.import_utils import LazyImportModule + + +if TYPE_CHECKING: + from .frame_sampler import (FRAME_SAMPLERS, IntervalSampler, SegmentSampler, + UniformSampler, do_frame_sample) + from .video_reader import (EasyVideoReader, FramesReaderWrapper, + VideoReaderWrapper) +else: + _import_structure = { + 'frame_sampler': ['FRAME_SAMPLERS', 'IntervalSampler', 'SegmentSampler', + 'UniformSampler', 'do_frame_sample'], + 'video_reader': ['EasyVideoReader', 'FramesReaderWrapper', 'VideoReaderWrapper'] + } + + import sys + sys.modules[__name__] = LazyImportModule( + __name__, + globals()['__file__'], + _import_structure, + module_spec=__spec__, + extra_objects={}, + ) diff --git a/scepter/modules/utils/visualization.py b/scepter/modules/utils/visualization.py index 82ba90f..3255031 100644 --- a/scepter/modules/utils/visualization.py +++ b/scepter/modules/utils/visualization.py @@ -1,8 +1,8 @@ # -*- coding: utf-8 -*- # Copyright (c) Alibaba, Inc. and its affiliates. import copy -from enum import Enum import os +from enum import Enum from scepter.modules.utils.file_system import FS @@ -17,17 +17,15 @@ class Media(Enum): class HtmlVisualization(object): - def __init__( - self, - allow_annotation=False, - slice_size=1000, - align='center', - width_scale='60%', - title='Visualization', - height=600, - width=None, - text_cols=40 - ): + def __init__(self, + allow_annotation=False, + slice_size=1000, + align='center', + width_scale='60%', + title='Visualization', + height=600, + width=None, + text_cols=40): self.content_list = [] self.rows_meta = [] self.allow_annotation = allow_annotation @@ -37,9 +35,9 @@ class HtmlVisualization(object): self.title = title self.html_start = '' self.html_head = f'{title}' - self.height = height if height is not None else "600" - self.width = width if width is not None else "auto" - self.text_cols = text_cols if text_cols is not None else "auto" + self.height = height if height is not None else '600' + self.width = width if width is not None else 'auto' + self.text_cols = text_cols if text_cols is not None else 'auto' self.html_style = (''' \n \n - '''.replace('{width_scale}', - self.width_scale).replace('{align}', self.align) - .replace('{pair_height}', f'{self.height}')) + '''.replace('{width_scale}', self.width_scale).replace( + '{align}', self.align).replace('{pair_height}', f'{self.height}')) self.html_body_script = ''' \n @@ -165,17 +158,16 @@ class HtmlVisualization(object): ''' self.label_button = ( - '
' + - "" - + '
') + '
' + + "" + + '
') def format_col(self, content='', label='', type=Media.TEXT, show_label=True, - cols_span=1 - ): + cols_span=1): if type == Media.TEXT: ret_str = '"{content}"' - sec_ret_str = f'{label}' if show_label else "" + sec_ret_str = f'{label}' if show_label else '' elif type == Media.IMAGE: ret_str = f'{label}' if show_label else "" + sec_ret_str = f'{label}' if show_label else '' elif type == Media.VIDEO: ret_str = '' - sec_ret_str = f'{label}' if show_label else "" + sec_ret_str = f'{label}' if show_label else '' elif type == Media.AUDIO: ret_str = f'