Compare commits
18
Commits
test-hf-sync
...
dmd
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
eb16ffb7b8 | ||
|
|
acd65e42e4 | ||
|
|
892d13d14d | ||
|
|
3a93f954df | ||
|
|
3e5ff5b583 | ||
|
|
84a6c32c96 | ||
|
|
ec9dd96d2e | ||
|
|
ed74b3c65c | ||
|
|
fce56124d7 | ||
|
|
4109928d27 | ||
|
|
d4ca37df9e | ||
|
|
9099e88e9a | ||
|
|
b679c8e515 | ||
|
|
da6003bc50 | ||
|
|
e811130464 | ||
|
|
d8a45e71c1 | ||
|
|
7b50887e38 | ||
|
|
89add12d3c |
@@ -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
|
||||
)
|
||||
|
||||
@@ -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
@@ -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
@@ -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")
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -46,6 +46,7 @@ def prepare_latents(
|
||||
return latents
|
||||
|
||||
|
||||
|
||||
def sample_validation_video(
|
||||
transformer,
|
||||
vae,
|
||||
|
||||
@@ -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"
|
||||
@@ -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
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user