Compare commits

...
Author SHA1 Message Date
Hangliang Ding eb16ffb7b8 Create distill_hunyuan_dmd_exp2.sh 2025-02-07 01:17:33 +08:00
Hangliang Ding acd65e42e4 Update num_frames 2025-02-07 01:15:57 +08:00
foreverpiano 892d13d14d fix on runpod 2025-01-11 14:13:09 +00:00
foreverpiano 3a93f954df update 2025-01-07 06:20:12 +00:00
foreverpiano 3e5ff5b583 update hardcode 2025-01-07 06:03:58 +00:00
foreverpiano 84a6c32c96 fix bugs 2025-01-05 09:03:10 +00:00
foreverpiano ec9dd96d2e update hunyuan runable version + freeze 2025-01-05 08:31:52 +00:00
foreverpiano ed74b3c65c update 2025-01-04 14:35:06 +00:00
foreverpiano fce56124d7 update predict_noise G 2025-01-04 12:00:23 +00:00
foreverpiano 4109928d27 fix some bug for hunyuan 2025-01-02 13:28:45 +00:00
“BrianChen1129” d4ca37df9e update 2025-01-01 04:35:36 +00:00
“BrianChen1129” 9099e88e9a 4 card adv 2025-01-01 04:27:35 +00:00
“BrianChen1129” b679c8e515 adv 2025-01-01 03:57:10 +00:00
foreverpiano da6003bc50 update disc 2024-12-29 09:57:37 +00:00
foreverpiano e811130464 update 2024-12-29 09:20:21 +00:00
foreverpiano d8a45e71c1 update 2024-12-27 14:35:12 +00:00
foreverpiano 7b50887e38 update 2024-12-27 14:00:07 +00:00
foreverpiano 89add12d3c update 2024-12-26 15:31:29 +00:00
13 changed files with 1656 additions and 73 deletions
+1 -1
View File
@@ -107,7 +107,7 @@ def latent_collate_function(batch):
if __name__ == "__main__":
dataset = LatentDataset("data/Mochi-Synthetic-Data/merge.txt", num_latent_t=28)
dataset = LatentDataset("data/HD-Mixkit-Finetune-Hunyuan/videos2caption.json", num_latent_t=8, cfg_rate=6)
dataloader = torch.utils.data.DataLoader(
dataset, batch_size=2, shuffle=False, collate_fn=latent_collate_function
)
+44 -6
View File
@@ -24,7 +24,7 @@ logger = logging.get_logger(__name__) # pylint: disable=invalid-name
class DiscriminatorHead(nn.Module):
def __init__(self, input_channel, output_channel=1):
def __init__(self, input_channel, output_channel=1, args=None):
super().__init__()
inner_channel = 1024
self.conv1 = nn.Sequential(
@@ -43,13 +43,21 @@ class DiscriminatorHead(nn.Module):
)
self.conv_out = nn.Conv2d(inner_channel, output_channel, 1, 1, 0)
vae_spatial_scale_factor = 8
self.patch_height = args.num_height // vae_spatial_scale_factor // 2
self.patch_width = args.num_width // vae_spatial_scale_factor // 2
print("## DiscriminatorHead: patch_height: ", self.patch_height)
print("## DiscriminatorHead: patch_width: ", self.patch_width)
def forward(self, x):
b, twh, c = x.shape
t = twh // (30 * 53)
x = x.view(-1, 30 * 53, c)
t = twh // (self.patch_height * self.patch_width)
x = x.view(-1, self.patch_height * self.patch_width, c)
x = x.permute(0, 2, 1)
x = x.view(b * t, c, 30, 53)
x = x.view(b * t, c, self.patch_height, self.patch_width)
x = self.conv1(x)
x = self.conv2(x) + x
x = self.conv_out(x)
@@ -58,7 +66,12 @@ class DiscriminatorHead(nn.Module):
class Discriminator(nn.Module):
def __init__(
self, stride=8, num_h_per_head=1, adapter_channel_dims=[3072], total_layers=48,
self,
stride=8,
num_h_per_head=1,
adapter_channel_dims=[3072],
total_layers = 48,
args=None,
):
super().__init__()
adapter_channel_dims = adapter_channel_dims * (total_layers // stride)
@@ -69,7 +82,10 @@ class Discriminator(nn.Module):
[
nn.ModuleList(
[
DiscriminatorHead(adapter_channel)
DiscriminatorHead(
adapter_channel,
args=args
)
for _ in range(self.num_h_per_head)
]
)
@@ -97,3 +113,25 @@ class Discriminator(nn.Module):
out = h(features[i])
outputs.append(out)
return outputs
class DMDiscriminator(nn.Module):
def __init__(self):
super().__init__()
self.cls_pred_branch = nn.Sequential(
nn.Conv2d(kernel_size=4, in_channels=1280, out_channels=1280, stride=2, padding=1), # 8x8 -> 4x4
nn.GroupNorm(num_groups=32, num_channels=1280),
nn.SiLU(),
nn.Conv2d(kernel_size=4, in_channels=1280, out_channels=1280, stride=4, padding=0), # 4x4 -> 1x1
nn.GroupNorm(num_groups=32, num_channels=1280),
nn.SiLU(),
nn.Conv2d(kernel_size=1, in_channels=1280, out_channels=1, stride=1, padding=0), # 1x1 -> 1x1
)
self.cls_pred_branch.requires_grad_(True)
def forward(self, features):
print("## features shape: ", features.shape)
return self.cls_pred_branch(features)
+72 -24
View File
@@ -22,7 +22,7 @@ from torch.distributed.fsdp import (
StateDictType,
FullStateDictConfig,
)
from fastvideo.utils.load import load_transformer
from fastvideo.utils.load import load_transformer
from fastvideo.models.mochi_hf.pipeline_mochi import linear_quadratic_schedule
import json
@@ -37,7 +37,9 @@ from fastvideo.utils.fsdp_util import (
get_discriminator_fsdp_kwargs,
)
import diffusers
from diffusers import FlowMatchEulerDiscreteScheduler
from diffusers import (
FlowMatchEulerDiscreteScheduler,
)
from fastvideo.distill.discriminator import Discriminator
from fastvideo.distill.solver import EulerSolver, extract_into_tensor
from copy import deepcopy
@@ -47,16 +49,45 @@ from diffusers.utils import check_min_version
from fastvideo.dataset.latent_datasets import LatentDataset, latent_collate_function
import torch.distributed as dist
from peft import LoraConfig
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
from torch.distributed.fsdp import (
FullyShardedDataParallel as FSDP,
)
from fastvideo.utils.checkpoint import (
save_checkpoint,
save_lora_checkpoint,
resume_lora_optimizer,
resume_training,
save_checkpoint_generator_discriminator,
resume_training_generator_discriminator,
)
# from fastvideo.utils.checkpoint import save_checkpoint
from fastvideo.utils.logging_ import main_print
from torch.distributed.fsdp import FullOptimStateDictConfig
from safetensors.torch import save_file
def save_checkpoint(model, rank, output_dir, step, discriminator=False):
with FSDP.state_dict_type(
model,
StateDictType.FULL_STATE_DICT,
FullStateDictConfig(offload_to_cpu=True, rank0_only=True),
FullOptimStateDictConfig(offload_to_cpu=True, rank0_only=True),
):
cpu_state = model.state_dict()
# todo move to get_state_dict
save_dir = os.path.join(output_dir, f"checkpoint-{step}")
os.makedirs(save_dir, exist_ok=True)
# save using safetensors
if rank <= 0 and not discriminator:
weight_path = os.path.join(save_dir, "diffusion_pytorch_model.safetensors")
save_file(cpu_state, weight_path)
config_dict = dict(model.config)
config_path = os.path.join(save_dir, "config.json")
# save dict as json
with open(config_path, "w") as f:
json.dump(config_dict, f, indent=4)
else:
weight_path = os.path.join(save_dir, "discriminator_pytorch_model.safetensors")
save_file(cpu_state, weight_path)
# Will error if the minimal version of diffusers is not installed. Remove at your own risks.
check_min_version("0.31.0")
@@ -73,7 +104,7 @@ def gan_d_loss(
encoder_hidden_states,
encoder_attention_mask,
weight,
discriminator_head_stride,
discriminator_head_stride
):
loss = 0.0
# collate sample_fake and sample_real
@@ -115,7 +146,7 @@ def gan_g_loss(
encoder_hidden_states,
encoder_attention_mask,
weight,
discriminator_head_stride,
discriminator_head_stride
):
loss = 0.0
features = teacher_transformer(
@@ -127,7 +158,9 @@ def gan_g_loss(
output_features_stride=discriminator_head_stride,
return_dict=False,
)[1]
fake_outputs = discriminator(features,)
fake_outputs = discriminator(
features,
)
for fake_output in fake_outputs:
loss += torch.mean(weight * torch.relu(1 - fake_output.float())) / (
discriminator.head_num * discriminator.num_h_per_head
@@ -156,7 +189,7 @@ def distill_one_step_adv(
not_apply_cfg_solver,
distill_cfg,
adv_weight,
discriminator_head_stride,
discriminator_head_stride
):
optimizer.zero_grad()
discriminator_optimizer.zero_grad()
@@ -272,7 +305,7 @@ def distill_one_step_adv(
huber_c = 0.001
g_loss = torch.mean(
torch.sqrt((model_pred.float() - target.float()) ** 2 + huber_c ** 2) - huber_c
torch.sqrt((model_pred.float() - target.float()) ** 2 + huber_c**2) - huber_c
)
discriminator.requires_grad_(False)
with torch.autocast("cuda", dtype=torch.bfloat16):
@@ -284,7 +317,7 @@ def distill_one_step_adv(
encoder_hidden_states.float(),
encoder_attention_mask,
1.0,
discriminator_head_stride,
discriminator_head_stride
)
g_loss += g_gan_loss
g_loss.backward()
@@ -356,10 +389,7 @@ def main(args):
torch.float32 if args.master_weight_type == "fp32" else torch.bfloat16,
)
teacher_transformer = deepcopy(transformer)
discriminator = Discriminator(
args.discriminator_head_stride,
total_layers=48 if args.model_type == "mochi" else 40,
)
discriminator = Discriminator(args.discriminator_head_stride, total_layers = 48 if args.model_type =="mochi" else 40)
if args.use_lora:
transformer.requires_grad_(False)
@@ -397,9 +427,18 @@ def main(args):
transformer._no_split_modules = no_split_modules
fsdp_kwargs["auto_wrap_policy"] = fsdp_kwargs["auto_wrap_policy"](transformer)
transformer = FSDP(transformer, **fsdp_kwargs,)
teacher_transformer = FSDP(teacher_transformer, **fsdp_kwargs,)
discriminator = FSDP(discriminator, **discriminator_fsdp_kwargs,)
transformer = FSDP(
transformer,
**fsdp_kwargs,
)
teacher_transformer = FSDP(
teacher_transformer,
**fsdp_kwargs,
)
discriminator = FSDP(
discriminator,
**discriminator_fsdp_kwargs,
)
main_print(f"--> model loaded")
if args.gradient_checkpointing:
@@ -570,7 +609,6 @@ def main(args):
if step <= int(phase_step):
return int(phase)
return phase
for i in range(init_steps):
_ = next(loader)
for step in range(init_steps + 1, args.max_train_steps + 1):
@@ -603,7 +641,7 @@ def main(args):
args.not_apply_cfg_solver,
args.distill_cfg,
args.adv_weight,
args.discriminator_head_stride,
args.discriminator_head_stride
)
step_time = time.time() - start_time
@@ -642,7 +680,7 @@ def main(args):
)
else:
# Your existing checkpoint saving code
# TODO
# TODO
# save_checkpoint_generator_discriminator(
# transformer,
# optimizer,
@@ -652,9 +690,7 @@ def main(args):
# args.output_dir,
# step,
# )
save_checkpoint(
transformer, rank, args.output_dir, args.max_train_steps
)
save_checkpoint(transformer, rank, args.output_dir, step, discriminator)
main_print(f"--> checkpoint saved at step {step}")
dist.barrier()
if args.log_validation and step % args.validation_steps == 0:
@@ -671,6 +707,7 @@ def main(args):
linear_range=args.linear_range,
ema=False,
)
if args.use_lora:
save_lora_checkpoint(
@@ -691,6 +728,8 @@ if __name__ == "__main__":
)
# dataset & dataloader
parser.add_argument("--data_json_path", type=str, required=True)
parser.add_argument("--num_height", type=int, default=480)
parser.add_argument("--num_width", type=int, default=848)
parser.add_argument("--num_frames", type=int, default=163)
parser.add_argument(
"--dataloader_num_workers",
@@ -928,6 +967,15 @@ if __name__ == "__main__":
default=0.025,
help="The threshold of the linear quadratic scheduler.",
)
parser.add_argument(
"--linear_range",
type=float,
default=0.5,
help="Range for linear quadratic scheduler.",
)
parser.add_argument(
"--weight_decay", type=float, default=0.001, help="Weight decay to apply."
)
parser.add_argument(
"--master_weight_type",
type=str,
@@ -935,4 +983,4 @@ if __name__ == "__main__":
help="Weight type to use - fp32 or bf16.",
)
args = parser.parse_args()
main(args)
main(args)
File diff suppressed because it is too large Load Diff
+31 -1
View File
@@ -49,7 +49,7 @@ def save_checkpoint(model, optimizer, rank, output_dir, step, discriminator=Fals
save_file(cpu_state, weight_path)
optimizer_path = os.path.join(save_dir, "discriminator_optimizer.pt")
torch.save(optim_state, optimizer_path)
def save_checkpoint_generator_discriminator(
model, optimizer, discriminator, discriminator_optimizer, rank, output_dir, step,
@@ -183,6 +183,36 @@ def resume_training_generator_discriminator(
)
return model, optimizer, discriminator, discriminator_optimizer, step
def resume_training_generator_fake_transformer(
model,
optimizer,
fake_transformer,
guidance_optimizer,
discriminator,
discriminator_optimizer,
checkpoint_dir,
rank,
):
step = int(checkpoint_dir.split("-")[-1])
model_weight_dir = os.path.join(checkpoint_dir, "model_weights_state")
model_optimizer_dir = os.path.join(checkpoint_dir, "model_optimizer_state")
fake_model_weight_dir = os.path.join(checkpoint_dir, "fake_model_weights_state")
fake_model_optimizer_dir = os.path.join(checkpoint_dir, "fake_model_optimizer_state")
model, optimizer = load_sharded_model(
model, optimizer, model_weight_dir, model_optimizer_dir
)
fake_transformer, guidance_optimizer = load_sharded_model(
fake_transformer, guidance_optimizer, fake_model_weight_dir, fake_model_optimizer_dir
)
discriminator_ckpt_file = os.path.join(
checkpoint_dir, "discriminator_fsdp_state", "discriminator_state.pt"
)
discriminator, discriminator_optimizer = load_full_state_model(
discriminator, discriminator_optimizer, discriminator_ckpt_file, rank
)
return model, optimizer, fake_transformer, guidance_optimizer, discriminator, discriminator_optimizer, step
def resume_training(model, optimizer, checkpoint_dir, discriminator=False):
weight_path = os.path.join(checkpoint_dir, "diffusion_pytorch_model.safetensors")
+2
View File
@@ -274,6 +274,8 @@ def load_transformer(
in_channels=16, out_channels=16, **hunyuan_config, dtype=master_weight_type,
)
transformer = load_hunyuan_state_dict(transformer, dit_model_name_or_path)
if master_weight_type == torch.bfloat16:
transformer = transformer.bfloat16()
else:
raise ValueError(f"Unsupported model type: {model_type}")
return transformer
+20 -41
View File
@@ -1,20 +1,23 @@
from accelerate.logging import get_logger
import torch
logger = get_logger(__name__)
def get_optimizer(args, params_to_optimize, use_deepspeed: bool = False):
def get_optimizer(
params_to_optimize,
args,
lr=1e-5,
betas=(0.9, 0.999),
weight_decay=1e-3,
eps=1e-8,
):
# Optimizer creation
supported_optimizers = ["adam", "adamw", "prodigy"]
supported_optimizers = ["adam", "adamw"]
if args.optimizer not in supported_optimizers:
logger.warning(
print(
f"Unsupported choice of optimizer: {args.optimizer}. Supported optimizers include {supported_optimizers}. Defaulting to AdamW"
)
args.optimizer = "adamw"
if args.use_8bit_adam and not (args.optimizer.lower() not in ["adam", "adamw"]):
logger.warning(
print(
f"use_8bit_adam is ignored when optimizer is not set to 'Adam' or 'AdamW'. Optimizer was "
f"set to {args.optimizer.lower()}"
)
@@ -34,44 +37,20 @@ def get_optimizer(args, params_to_optimize, use_deepspeed: bool = False):
optimizer = optimizer_class(
params_to_optimize,
betas=(args.adam_beta1, args.adam_beta2),
eps=args.adam_epsilon,
weight_decay=args.adam_weight_decay,
lr=lr,
betas=betas,
eps=eps,
weight_decay=weight_decay,
)
elif args.optimizer.lower() == "adam":
optimizer_class = bnb.optim.Adam8bit if args.use_8bit_adam else torch.optim.Adam
optimizer = optimizer_class(
params_to_optimize,
betas=(args.adam_beta1, args.adam_beta2),
eps=args.adam_epsilon,
weight_decay=args.adam_weight_decay,
)
elif args.optimizer.lower() == "prodigy":
try:
import prodigyopt
except ImportError:
raise ImportError(
"To use Prodigy, please install the prodigyopt library: `pip install prodigyopt`"
)
optimizer_class = prodigyopt.Prodigy
if args.learning_rate <= 0.1:
logger.warning(
"Learning rate is too low. When using prodigy, it's generally better to set learning rate around 1.0"
)
optimizer = optimizer_class(
params_to_optimize,
lr=args.learning_rate,
betas=(args.adam_beta1, args.adam_beta2),
beta3=args.prodigy_beta3,
weight_decay=args.adam_weight_decay,
eps=args.adam_epsilon,
decouple=args.prodigy_decouple,
use_bias_correction=args.prodigy_use_bias_correction,
safeguard_warmup=args.prodigy_safeguard_warmup,
lr=lr,
betas=betas,
eps=eps,
weight_decay=weight_decay,
)
return optimizer
+1
View File
@@ -46,6 +46,7 @@ def prepare_latents(
return latents
def sample_validation_video(
transformer,
vae,
+44
View File
@@ -0,0 +1,44 @@
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
export WANDB_API_KEY=4f6de3765d6464f43e0506ec7d785641af645e73
DATA_DIR=./data
torchrun --nnodes 1 --nproc_per_node 4\
fastvideo/distill_adv.py\
--seed 42\
--pretrained_model_name_or_path $DATA_DIR/hunyuan\
--dit_model_name_or_path $DATA_DIR/hunyuan/hunyuan-video-t2v-720p/transformers/mp_rank_00_model_states.pt\
--model_type "hunyuan" \
--cache_dir "$DATA_DIR/.cache"\
--data_json_path "$DATA_DIR/HD-Mixkit-Finetune-Hunyuan/videos2caption.json"\
--validation_prompt_dir "$DATA_DIR/HD-Mixkit-Finetune-Hunyuan/validation"\
--gradient_checkpointing\
--train_batch_size=1\
--num_latent_t 8\
--sp_size 4 \
--train_sp_batch_size 1\
--dataloader_num_workers 4\
--gradient_accumulation_steps=1\
--max_train_steps=2000\
--learning_rate=1e-6\
--mixed_precision="bf16"\
--checkpointing_steps=64\
--validation_steps 64\
--validation_sampling_steps "2,4,8" \
--checkpoints_total_limit 3\
--allow_tf32\
--ema_start_step 0\
--cfg 0.0\
--log_validation\
--output_dir="$DATA_DIR/outputs/hy_phase1_shift17_bs_8_adv_HD"\
--tracker_project_name Hunyuan_Distill \
--num_height 720 \
--num_width 1280 \
--num_frames 125 \
--shift 17 \
--validation_guidance_scale "1.0" \
--num_euler_timesteps 50 \
--multi_phased_distill_schedule "4000-1" \
--not_apply_cfg_solver \
--master_weight_type "bf16"
+48
View File
@@ -0,0 +1,48 @@
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
export WANDB_API_KEY=7abd9763ff10869f9e88526d919c2639766139a8
DATA_DIR=./data
torchrun --nnodes 1 --nproc_per_node 4\
fastvideo/distill_dmd.py\
--seed 42\
--pretrained_model_name_or_path $DATA_DIR/hunyuan\
--dit_model_name_or_path $DATA_DIR/hunyuan/hunyuan-video-t2v-720p/transformers/mp_rank_00_model_states.pt\
--model_type "hunyuan" \
--cache_dir "$DATA_DIR/.cache"\
--data_json_path "$DATA_DIR/HD-Mixkit-Finetune-Hunyuan/videos2caption.json"\
--validation_prompt_dir "$DATA_DIR/HD-Mixkit-Finetune-Hunyuan/validation"\
--gradient_checkpointing\
--train_batch_size=1\
--num_latent_t 8\
--sp_size 4\
--train_sp_batch_size 1\
--dataloader_num_workers 4\
--gradient_accumulation_steps=1\
--max_train_steps=30000\
--learning_rate=1e-5\
--mixed_precision="bf16"\
--checkpointing_steps=64\
--validation_steps 64\
--validation_sampling_steps "2,4,8" \
--num_inference_steps 4 \
--checkpoints_total_limit 3\
--allow_tf32\
--ema_start_step 0\
--cfg 0.0\
--log_validation\
--output_dir="$DATA_DIR/outputs/hy_phase1_shift17_bs_32"\
--tracker_project_name Hunyuan_Distill \
--num_height 720 \
--num_width 1280 \
--num_frames 125 \
--shift 17 \
--validation_guidance_scale "1.0" \
--num_euler_timesteps 50 \
--multi_phased_distill_schedule "4000-1" \
--not_apply_cfg_solver \
--master_weight_type "bf16" \
--optimizer "AdamW" \
--use_8bit_adam \
--generator_update_steps 3
@@ -0,0 +1,48 @@
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
export WANDB_API_KEY=7abd9763ff10869f9e88526d919c2639766139a8
DATA_DIR=./data
torchrun --nnodes 1 --nproc_per_node 4\
fastvideo/distill_dmd.py\
--seed 42\
--pretrained_model_name_or_path $DATA_DIR/hunyuan\
--dit_model_name_or_path $DATA_DIR/hunyuan/hunyuan-video-t2v-720p/transformers/mp_rank_00_model_states.pt\
--model_type "hunyuan" \
--cache_dir "$DATA_DIR/.cache"\
--data_json_path "$DATA_DIR/HD-Mixkit-Finetune-Hunyuan/videos2caption.json"\
--validation_prompt_dir "$DATA_DIR/HD-Mixkit-Finetune-Hunyuan/validation"\
--gradient_checkpointing\
--train_batch_size=1\
--num_latent_t 8\
--sp_size 4\
--train_sp_batch_size 1\
--dataloader_num_workers 4\
--gradient_accumulation_steps=1\
--max_train_steps=30000\
--learning_rate=1e-5\
--mixed_precision="bf16"\
--checkpointing_steps=64\
--validation_steps 64\
--validation_sampling_steps "2,4,8" \
--num_inference_steps 8 \
--checkpoints_total_limit 3\
--allow_tf32\
--ema_start_step 0\
--cfg 0.0\
--log_validation\
--output_dir="$DATA_DIR/outputs/hy_phase1_shift17_bs_32"\
--tracker_project_name Hunyuan_Distill \
--num_height 720 \
--num_width 1280 \
--num_frames 125 \
--shift 17 \
--validation_guidance_scale "1.0" \
--num_euler_timesteps 50 \
--multi_phased_distill_schedule "4000-1" \
--not_apply_cfg_solver \
--master_weight_type "bf16" \
--optimizer "AdamW" \
--use_8bit_adam \
--generator_update_steps 3
@@ -0,0 +1,49 @@
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
export WANDB_API_KEY=7abd9763ff10869f9e88526d919c2639766139a8
DATA_DIR=./data
torchrun --nnodes 1 --nproc_per_node 4\
fastvideo/distill_dmd.py\
--seed 42\
--pretrained_model_name_or_path $DATA_DIR/hunyuan\
--dit_model_name_or_path $DATA_DIR/hunyuan/hunyuan-video-t2v-720p/transformers/mp_rank_00_model_states.pt\
--model_type "hunyuan" \
--cache_dir "$DATA_DIR/.cache"\
--data_json_path "$DATA_DIR/HD-Mixkit-Finetune-Hunyuan/videos2caption.json"\
--validation_prompt_dir "$DATA_DIR/HD-Mixkit-Finetune-Hunyuan/validation"\
--gradient_checkpointing\
--train_batch_size=1\
--num_latent_t 8\
--sp_size 4\
--train_sp_batch_size 1\
--dataloader_num_workers 4\
--gradient_accumulation_steps=1\
--max_train_steps=30000\
--learning_rate=1e-5\
--mixed_precision="bf16"\
--checkpointing_steps=1\
--validation_steps 1\
--validation_sampling_steps "2,4,8" \
--num_inference_steps 4 \
--checkpoints_total_limit 3\
--allow_tf32\
--ema_start_step 0\
--cfg 0.0\
--log_validation\
--output_dir="$DATA_DIR/outputs/hy_phase1_shift17_bs_32"\
--tracker_project_name Hunyuan_Distill \
--num_height 784 \
--num_width 1280 \
--num_frames 29 \
--shift 17 \
--validation_guidance_scale "1.0" \
--num_euler_timesteps 50 \
--multi_phased_distill_schedule "4000-1" \
--not_apply_cfg_solver \
--master_weight_type "bf16" \
--optimizer "AdamW" \
--use_8bit_adam \
--generator_update_steps 5 \
--run_pod
+39
View File
@@ -0,0 +1,39 @@
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
torchrun --nnodes 1 --nproc_per_node 4 \
fastvideo/distill_dmd.py \
--seed 42 \
--pretrained_model_name_or_path data/mochi \
--model_type "mochi" \
--cache_dir data/.cache \
--data_json_path data/Mochi-425-Data/videos2caption.json \
--validation_prompt_dir data/Mochi-425-Data/validation \
--gradient_checkpointing \
--train_batch_size=1 \
--num_latent_t 16 \
--sp_size 4 \
--train_sp_batch_size 1 \
--generator_update_steps=1 \
--dataloader_num_workers 4 \
--gradient_accumulation_steps=1 \
--max_train_steps=4000 \
--learning_rate=1e-6 \
--mixed_precision=bf16 \
--checkpointing_steps=64 \
--validation_steps=64 \
--validation_sampling_steps 6 \
--checkpoints_total_limit 3 \
--allow_tf32 \
--ema_start_step 0 \
--cfg 0.0 \
--log_validation \
--output_dir="data/outputs/lq_euler_50_thres0.1_lrg_0.75_phase1_lr1e-6_repro" \
--tracker_project_name PCM \
--num_frames 93 \
--scheduler_type pcm_linear_quadratic \
--validation_guidance_scale 0.5,1.5,2.5 \
--num_euler_timesteps 50 \
--linear_quadratic_threshold 0.1 \
--linear_range 0.75 \
--multi_phased_distill_schedule 4000-1