diff --git a/README.md b/README.md index 9418b9e..918360f 100644 --- a/README.md +++ b/README.md @@ -517,12 +517,11 @@ The required fields include the following six, with their explanations as follow All parameters related to training are stored in 'train_config/ace_plus_lora.yaml'. With the following default configuration, the memory usage for LoRA training is between 38GB and 40GB. -| Hyperparameter | Value | Description| -|::|:---:| +| Hyperparameter | Value | Description | +| --- | --- | --- | | ATTN_BACKEND | flash_attn / pytorch |Set 'flash_attn' to use flash_attn2(Make sure you have installed flash-attn2 correctly). If the version of PyTorch is greater than 2.4.0, use 'pytorch' to utilize PyTorch's implementation.| | USE_GRAD_CHECKPOINT | True / False |Using gradient checkpointing can also significantly reduce GPU memory usage, but it may slow down the training speed. | -| MAX_SEQ_LEN | 2048 | The MAX_SEQ_LEN refers to the sequence size limit for a single input image (calculated as H/16 * W/16). - A larger value indicates a longer computation sequence and a higher training resolution. The default value I provided is 2048.| +| MAX_SEQ_LEN | 2048 | The MAX_SEQ_LEN refers to the sequence size limit for a single input image (calculated as H/16 * W/16). A larger value indicates a longer computation sequence and a higher training resolution. The default value I provided is 2048.| To run the training code, execute the following command. diff --git a/requirements.txt b/requirements.txt index 8b13789..e69de29 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1 +0,0 @@ - diff --git a/train_config/ace_plus_fft.yaml b/train_config/ace_plus_fft.yaml index 5740324..b7f4b67 100644 --- a/train_config/ace_plus_fft.yaml +++ b/train_config/ace_plus_fft.yaml @@ -108,7 +108,10 @@ SOLVER: DEPTH: 19 # DEPTH_SINGLE_BLOCKS DESCRIPTION: number of transformer blocks in the single stream block. TYPE: int default: 38 DEPTH_SINGLE_BLOCKS: 38 + # ATTN_BACKEND:setting 'flash_attn' to use flash_attn2, if the version of pytorch > 2.4.0, using 'pytorch' to use pytorch's implementation ATTN_BACKEND: flash_attn + # USE_GRAD_CHECKPOINT: setting gc to true can decrease the memory usage. + USE_GRAD_CHECKPOINT: True # FIRST_STAGE_MODEL: diff --git a/train_config/ace_plus_fft_lora.yaml b/train_config/ace_plus_fft_lora.yaml index 226717e..3e5fe5e 100644 --- a/train_config/ace_plus_fft_lora.yaml +++ b/train_config/ace_plus_fft_lora.yaml @@ -108,7 +108,10 @@ SOLVER: DEPTH: 19 # DEPTH_SINGLE_BLOCKS DESCRIPTION: number of transformer blocks in the single stream block. TYPE: int default: 38 DEPTH_SINGLE_BLOCKS: 38 + # ATTN_BACKEND:setting 'flash_attn' to use flash_attn2, if the version of pytorch > 2.4.0, using 'pytorch' to use pytorch's implementation ATTN_BACKEND: flash_attn + # USE_GRAD_CHECKPOINT: setting gc to true can decrease the memory usage. + USE_GRAD_CHECKPOINT: True # FIRST_STAGE_MODEL: diff --git a/train_config/ace_plus_lora.yaml b/train_config/ace_plus_lora.yaml index dc261b2..4d19db2 100644 --- a/train_config/ace_plus_lora.yaml +++ b/train_config/ace_plus_lora.yaml @@ -108,7 +108,10 @@ SOLVER: DEPTH: 19 # DEPTH_SINGLE_BLOCKS DESCRIPTION: number of transformer blocks in the single stream block. TYPE: int default: 38 DEPTH_SINGLE_BLOCKS: 38 + # ATTN_BACKEND:setting 'flash_attn' to use flash_attn2, if the version of pytorch > 2.4.0, using 'pytorch' to use pytorch's implementation ATTN_BACKEND: flash_attn + # USE_GRAD_CHECKPOINT: setting gc to true can decrease the memory usage. + USE_GRAD_CHECKPOINT: True # FIRST_STAGE_MODEL: