Compare commits

..
Author SHA1 Message Date
JerryZhou54 bf5726cd06 Change dataloader 2025-05-28 00:17:46 +00:00
“BrianChen1129” 57cfd16136 update preprocess 2025-05-27 03:36:46 +00:00
9 changed files with 185 additions and 283 deletions
+161 -208
View File
@@ -1,21 +1,29 @@
import argparse
import json
import os
import random
import time
from collections import defaultdict
import numpy as np
import pyarrow.parquet as pq
import torch
import tqdm
from einops import rearrange
from torch import distributed as dist
from torch.utils.data import IterableDataset, get_worker_info
from torch.utils.data import Dataset
from torchdata.stateful_dataloader import StatefulDataLoader
from fastvideo.v1.distributed import (get_sequence_model_parallel_rank,
get_sp_group)
from fastvideo.v1.logger import init_logger
# Path to your dataset
dataset_path = "/mnt/sharefs/users/hao.zhang/Vchitect-2M/Vchitect-2M-laten-93x512x512/train/"
logger = init_logger(__name__)
class ParquetVideoTextDataset(IterableDataset):
class ParquetVideoTextDataset(Dataset):
"""Efficient loader for video-text data from a directory of Parquet files."""
def __init__(self,
@@ -24,237 +32,185 @@ class ParquetVideoTextDataset(IterableDataset):
rank: int = 0,
world_size: int = 1,
cfg_rate: float = 0.0,
num_latent_t: int = 2):
num_latent_t: int = 2,
seed: int = 0):
super().__init__()
self.path = str(path)
self.batch_size = batch_size
self.rank = rank
self.world_size = world_size
self.local_rank = get_sequence_model_parallel_rank()
self.sp_world_size = world_size
self.world_size = int(os.getenv("WORLD_SIZE", 1))
self.cfg_rate = cfg_rate
self.num_latent_t = num_latent_t
self.local_indices = None
self.plan_output_dir = os.path.join(self.path, "data_plan.json")
# Find all parquet files recursively
print(f"Scanning for parquet files in {self.path}")
self.parquet_files = []
for root, _, files in os.walk(self.path):
for file in files:
if file.endswith('.parquet'):
self.parquet_files.append(os.path.join(root, file))
# Sort files for consistent ordering
self.parquet_files.sort()
ranks = get_sp_group().ranks
group_ranks = [None for _ in range(self.world_size)]
torch.distributed.all_gather_object(group_ranks, ranks)
# Distribute files among workers
# drop last unenven files
print(f"Total files: {len(self.parquet_files)}")
total_files = len(self.parquet_files)
base_count = total_files // world_size
extra_files = total_files % world_size
if rank == 0:
# If a plan already exists, then skip creating a new plan
# This will be useful when resume training
if os.path.exists(self.plan_output_dir):
print(f"Using existing plan from {self.plan_output_dir}")
return
if rank < extra_files:
start_idx = rank * (base_count + 1)
end_idx = start_idx + base_count + 1
else:
start_idx = rank * base_count + extra_files
end_idx = start_idx + base_count
# Find all parquet files recursively, and record num_rows for each file
print(f"Scanning for parquet files in {self.path}")
metadatas = []
for root, _, files in os.walk(self.path):
for file in sorted(files):
if file.endswith('.parquet'):
file_path = os.path.join(root, file)
num_rows = pq.ParquetFile(file_path).metadata.num_rows
for row_idx in range(num_rows):
metadatas.append((file_path, row_idx))
self.parquet_files = self.parquet_files[start_idx:end_idx]
# Generate the plan that distribute rows among workers
random.seed(seed)
random.shuffle(metadatas)
print(f"Files assigned to rank {rank}: {len(self.parquet_files)}")
if len(self.parquet_files) > 0:
print(f"First file: {self.parquet_files[0]}")
print(f"Last file: {self.parquet_files[-1]}")
# Get all sp groups
# e.g. if num_gpus = 4, sp_size = 2
# group_ranks = [(0, 1), (2, 3)]
# We will assign the same batches of data to ranks in the same sp group, and we'll assign different batches to ranks in different sp groups
# e.g. plan = {0: [row 1, row 4], 1: [row 1, row 4], 2: [row 2, row 3], 3: [row 2, row 3]}
group_ranks = list(set(tuple(r) for r in group_ranks))
num_sp_groups = len(group_ranks)
plan = defaultdict(list)
for idx, metadata in enumerate(metadatas):
sp_group_idx = idx % num_sp_groups
for global_rank in group_ranks[sp_group_idx]:
plan[global_rank].append(metadata)
# Initialize current file index
self.current_file_idx = 0
self.current_reader = None
self.current_batches = None
self.total_samples = 0
with open(self.plan_output_dir, "w") as f:
json.dump(plan, f)
def _open_next_file(self):
"""Open the next parquet file for reading."""
num_workers = get_worker_info().num_workers
worker_id = get_worker_info().id
total_files = len(self.parquet_files)
base_count = total_files // num_workers
extra_files = total_files % num_workers
if worker_id < extra_files:
start_idx = worker_id * (base_count + 1)
end_idx = start_idx + base_count + 1
else:
start_idx = worker_id * base_count + extra_files
end_idx = start_idx + base_count
worker_parquet_files = self.parquet_files[start_idx:end_idx]
if self.current_file_idx >= len(worker_parquet_files):
print(
f"Rank {self.rank}, Worker {worker_id}: No more files to open (current_idx={self.current_file_idx}, total_files={len(worker_parquet_files)})"
)
return False
if self.current_reader is not None:
self.current_reader.close()
file_path = worker_parquet_files[self.current_file_idx]
print(
f"Rank {self.rank}, Worker {worker_id}: Opening file {self.current_file_idx + 1}/{len(worker_parquet_files)}: {file_path}"
)
try:
self.current_reader = pq.ParquetFile(file_path)
self.current_batches = self.current_reader.iter_batches(
batch_size=self.batch_size)
self.current_file_idx += 1
return True
except Exception as e:
print(f"Error opening file {file_path}: {str(e)}")
return False
def __iter__(self):
"""Iterate over the dataset in a streaming fashion."""
print(f"Rank {self.rank}: Starting iteration")
# First try to open a file
if not self._open_next_file():
print(f"Rank {self.rank}: Failed to open first file")
return
while True:
def __len__(self):
if self.local_indices is None:
try:
# Get next batch from current file
batch = next(self.current_batches)
batch_dict = batch.to_pydict()
processed = self._process_batch(batch_dict)
with open(self.plan_output_dir) as f:
plan = json.load(f)
self.local_indices = plan[str(self.rank)]
except:
raise Exception("The data plan hasn't been created yet")
return len(self.local_indices)
# Update sample count
batch_size = len(processed["latents"])
self.total_samples += batch_size
def __getitem__(self, idx):
if self.local_indices is None:
try:
with open(self.plan_output_dir) as f:
plan = json.load(f)
self.local_indices = plan[self.rank]
except:
raise Exception("The data plan hasn't been created yet")
file_path, row_idx = self.local_indices[idx]
parquet_file = pq.ParquetFile(file_path)
# Print progress
if self.total_samples % 1000 == 0:
print(
f"Rank {self.rank}: Processed {self.total_samples} samples"
)
# Calculate the row group to read into memory and the local idx
# This way we can avoid reading in the entire parquet file
cumulative = 0
for i in range(parquet_file.num_row_groups):
num_rows = parquet_file.metadata.row_group(i).num_rows
if cumulative + num_rows > idx:
row_group_index = i
local_index = idx - cumulative
break
cumulative += num_rows
# Yield each item in the batch
for lat, emb, mask, info in zip(processed["latents"],
processed["embeddings"],
processed["masks"],
processed["info"]):
if lat.numel() == 0: # Split is validation
yield lat, emb, mask, info
else:
yield lat[:, -self.num_latent_t:], emb, mask, info
row_group = parquet_file.read_row_group(row_group_index).to_pydict()
row_dict = {k: v[local_index] for k, v in row_group.items()}
del row_group
except StopIteration:
# Current file is exhausted, try next file
print(
f"Rank {self.rank}: Current file exhausted, trying next file"
)
self.current_batches = None
if not self._open_next_file():
print(
f"Rank {self.rank}: No more files to process. Total samples: {self.total_samples}"
)
break
except Exception as e:
print(f"Error processing batch: {str(e)}")
self.current_batches = None
if not self._open_next_file():
print(
f"Rank {self.rank}: Failed to open next file after error"
)
break
processed = self._process_row(row_dict)
lat, emb, mask, info = processed["latents"], processed[
"embeddings"], processed["masks"], processed["info"]
if lat.numel() == 0: # Validation parquet
return lat, emb, mask, info
else:
lat = lat[:, -self.num_latent_t:]
if self.sp_world_size > 1:
lat = rearrange(lat,
"t (n s) h w -> t n s h w",
n=self.sp_world_size).contiguous()
lat = lat[:, self.local_rank, :, :, :]
return lat, emb, mask, info
# Clean up
if self.current_reader is not None:
self.current_reader.close()
def _process_batch(self, batch):
def _process_row(self, row):
"""Process a PyArrow batch into tensors."""
out = {"lat": [], "emb": [], "msk": [], "info": []}
out = {"lat": None, "emb": None, "msk": None, "info": None}
for i in range(len(batch["vae_latent_bytes"])):
vae_latent_bytes = batch["vae_latent_bytes"][i]
vae_latent_shape = batch["vae_latent_shape"][i]
text_embedding_bytes = batch["text_embedding_bytes"][i]
text_embedding_shape = batch["text_embedding_shape"][i]
text_attention_mask_bytes = batch["text_attention_mask_bytes"][i]
text_attention_mask_shape = batch["text_attention_mask_shape"][i]
vae_latent_bytes = row["vae_latent_bytes"]
vae_latent_shape = row["vae_latent_shape"]
text_embedding_bytes = row["text_embedding_bytes"]
text_embedding_shape = row["text_embedding_shape"]
text_attention_mask_bytes = row["text_attention_mask_bytes"]
text_attention_mask_shape = row["text_attention_mask_shape"]
# Process latent
if not vae_latent_shape: # No VAE latent is stored. Split is validation
lat = np.array([])
else:
lat = np.frombuffer(vae_latent_bytes,
dtype=np.float32).reshape(vae_latent_shape)
# Make array writable
lat = np.copy(lat)
# Process latent
if not vae_latent_shape: # No VAE latent is stored. Split is validation
lat = np.array([])
else:
lat = np.frombuffer(vae_latent_bytes,
dtype=np.float32).reshape(vae_latent_shape)
# Make array writable
lat = np.copy(lat)
if random.random() < self.cfg_rate:
emb = np.zeros((512, 4096), dtype=np.float32)
else:
emb = np.frombuffer(
text_embedding_bytes,
dtype=np.float32).reshape(text_embedding_shape)
# Make array writable
emb = np.copy(emb)
if emb.shape[0] < 512:
padded_emb = np.zeros((512, emb.shape[1]), dtype=np.float32)
padded_emb[:emb.shape[0], :] = emb
emb = padded_emb
elif emb.shape[0] > 512:
emb = emb[:512, :]
if random.random() < self.cfg_rate:
emb = np.zeros((512, 4096), dtype=np.float32)
else:
emb = np.frombuffer(text_embedding_bytes,
dtype=np.float32).reshape(text_embedding_shape)
# Make array writable
emb = np.copy(emb)
if emb.shape[0] < 512:
padded_emb = np.zeros((512, emb.shape[1]), dtype=np.float32)
padded_emb[:emb.shape[0], :] = emb
emb = padded_emb
elif emb.shape[0] > 512:
emb = emb[:512, :]
# Process mask
if len(text_attention_mask_bytes) > 0 and len(
text_attention_mask_shape) > 0:
msk = np.frombuffer(text_attention_mask_bytes,
dtype=np.uint8).astype(np.bool_)
msk = msk.reshape(1, -1)
# Make array writable
msk = np.copy(msk)
if msk.shape[1] < 512:
padded_msk = np.zeros((1, 512), dtype=np.bool_)
padded_msk[:, :msk.shape[1]] = msk
msk = padded_msk
elif msk.shape[1] > 512:
msk = msk[:, :512]
else:
msk = np.ones((1, 512), dtype=np.bool_)
# to string
file_name = str(batch["file_name"][i])
# Collect metadata
info = {
"width": batch["width"][i],
"height": batch["height"][i],
"num_frames": batch["num_frames"][i],
"duration_sec": batch["duration_sec"][i],
"fps": batch["fps"][i],
"file_name": batch["file_name"][i],
"caption": batch["caption"][i],
}
# Process mask
if len(text_attention_mask_bytes) > 0 and len(
text_attention_mask_shape) > 0:
msk = np.frombuffer(text_attention_mask_bytes,
dtype=np.uint8).astype(np.bool_)
msk = msk.reshape(1, -1)
# Make array writable
msk = np.copy(msk)
if msk.shape[1] < 512:
padded_msk = np.zeros((1, 512), dtype=np.bool_)
padded_msk[:, :msk.shape[1]] = msk
msk = padded_msk
elif msk.shape[1] > 512:
msk = msk[:, :512]
else:
msk = np.ones((1, 512), dtype=np.bool_)
out["lat"].append(torch.from_numpy(lat))
out["emb"].append(torch.from_numpy(emb))
out["msk"].append(torch.from_numpy(msk))
out["info"].append(info)
return {
"latents": torch.stack(out["lat"]) if out["lat"] else None,
"embeddings": torch.stack(out["emb"]) if out["emb"] else None,
"masks": torch.stack(out["msk"]) if out["msk"] else None,
"info": out["info"]
# Collect metadata
info = {
"width": row["width"],
"height": row["height"],
"num_frames": row["num_frames"],
"duration_sec": row["duration_sec"],
"fps": row["fps"],
"file_name": row["file_name"],
"caption": row["caption"],
}
out["lat"] = torch.from_numpy(lat)
out["emb"] = torch.from_numpy(emb)
out["msk"] = torch.from_numpy(msk)
out["info"] = info
def bind_cpu_cores(local_rank, cpu_per_process=16):
"""根据local_rank绑定固定cpu核。"""
start = local_rank * cpu_per_process
end = start + cpu_per_process
cores = list(range(start, end))
print(f"[Rank {local_rank}] Binding to CPU cores: {cores}")
os.sched_setaffinity(0, cores)
return {
"latents": out["lat"],
"embeddings": out["emb"],
"masks": out["msk"],
"info": out["info"]
}
if __name__ == "__main__":
@@ -297,9 +253,6 @@ if __name__ == "__main__":
f"Initialized process: rank={rank}, local_rank={local_rank}, world_size={world_size}, device={device}"
)
# Bind CPU cores after distributed initialization
# bind_cpu_cores(local_rank, cpu_per_process=16)
# Create dataset
dataset = ParquetVideoTextDataset(
args.path,
+2 -16
View File
@@ -394,16 +394,7 @@ class TransformerLoader(ComponentLoader):
default_dtype = PRECISION_TO_TYPE[fastvideo_args.precision]
# Load the model using FSDP loader
logger.info("Loading model from %s, default_dtype: %s", cls_name, default_dtype)
# model = load_fsdp_model(model_cls=model_cls,
# init_params={
# "config": dit_config,
# "hf_config": hf_config
# },
# weight_dir_list=safetensors_list,
# device=fastvideo_args.device,
# cpu_offload=fastvideo_args.use_cpu_offload,
# default_dtype=default_dtype)
logger.info("Loading model from %s", cls_name)
model = load_fsdp_model(model_cls=model_cls,
init_params={
"config": dit_config,
@@ -412,12 +403,7 @@ class TransformerLoader(ComponentLoader):
weight_dir_list=safetensors_list,
device=fastvideo_args.device,
cpu_offload=fastvideo_args.use_cpu_offload,
default_dtype=default_dtype,
# TODO(will): make these configurable
param_dtype=torch.bfloat16,
reduce_dtype=torch.float32,
output_dtype=None,
)
default_dtype=default_dtype)
if fastvideo_args.enable_torch_compile:
logger.info("Torch Compile enabled for DiT")
for n, m in reversed(list(model.named_modules())):
+4 -29
View File
@@ -14,16 +14,13 @@ from typing import (Any, Callable, DefaultDict, Dict, Generator, Hashable, List,
import torch
from torch import nn
from torch.distributed import DeviceMesh, init_device_mesh
from torch.distributed.fsdp import CPUOffloadPolicy, fully_shard, MixedPrecisionPolicy
from torch.distributed._composable.fsdp import CPUOffloadPolicy, fully_shard
from torch.distributed._tensor import distribute_tensor
from torch.nn.modules.module import _IncompatibleKeys
from fastvideo.v1.distributed.parallel_state import (
get_sequence_model_parallel_world_size)
from fastvideo.v1.models.loader.weight_utils import safetensors_weights_iterator
from fastvideo.v1.logger import init_logger
logger = init_logger(__name__)
# TODO(PY): move this to utils elsewhere
@@ -89,29 +86,16 @@ def get_param_names_mapping(
# TODO(PY): add compile option
# param_dtype: torch.dtype,
# reduce_dtype: torch.dtype,
# output_dtype: torch.dtype,
# pp_enabled: bool = False,
# cpu_offload: bool = False,
def load_fsdp_model(
model_cls: Type[nn.Module],
init_params: Dict[str, Any],
weight_dir_list: List[str],
device: torch.device,
default_dtype: torch.dtype,
param_dtype: torch.dtype,
reduce_dtype: torch.dtype,
cpu_offload: bool = False,
output_dtype: Optional[torch.dtype] = None,
default_dtype: Optional[torch.dtype] = torch.bfloat16,
) -> torch.nn.Module:
mp_policy = MixedPrecisionPolicy(param_dtype, reduce_dtype, output_dtype, cast_forward_inputs=True)
# with set_default_dtype(default_dtype), torch.device("meta"):
with set_default_dtype(default_dtype), torch.device("meta"):
model = model_cls(**init_params)
device_mesh = init_device_mesh(
"cuda",
mesh_shape=(get_sequence_model_parallel_world_size(), ),
@@ -120,7 +104,6 @@ def load_fsdp_model(
shard_model(model,
cpu_offload=cpu_offload,
reshard_after_forward=True,
mp_policy=mp_policy,
dp_mesh=device_mesh["dp"])
weight_iterator = safetensors_weights_iterator(weight_dir_list)
param_names_mapping_fn = get_param_names_mapping(model._param_names_mapping)
@@ -147,7 +130,6 @@ def shard_model(
*,
cpu_offload: bool,
reshard_after_forward: bool = True,
mp_policy: Optional[MixedPrecisionPolicy] = None,
dp_mesh: Optional[DeviceMesh] = None,
) -> None:
"""
@@ -175,17 +157,14 @@ def shard_model(
"""
fsdp_kwargs = {
"reshard_after_forward": reshard_after_forward,
"mesh": dp_mesh,
"mp_policy": mp_policy,
"mesh": dp_mesh
}
if cpu_offload:
fsdp_kwargs["offload_policy"] = CPUOffloadPolicy()
# iterating in reverse to start with
# Shard the model with FSDP, iterating in reverse to start with
# lowest-level modules first
num_layers_sharded = 0
# TODO(will): don't reshard after forward for the last layer to save on the
# all-gather that will immediately happen Shard the model with FSDP,
for n, m in reversed(list(model.named_modules())):
if any([
shard_condition(n, m)
@@ -232,10 +211,6 @@ def load_fsdp_model_from_full_model_state_dict(
NotImplementedError: If got FSDP with more than 1D.
"""
meta_sharded_sd = model.state_dict()
# s = fully_shard.state(model)
# logger.info(f"type(s): {type(s)}")
# logger.info(f"s: {s}")
# import pdb; pdb.set_trace()
sharded_sd = {}
to_merge_params: DefaultDict[Hashable, Dict[Any, Any]] = defaultdict(dict)
@@ -155,21 +155,17 @@ class ComposedPipelineBase(ABC):
for key, value in config_args.items():
setattr(fastvideo_args, key, value)
# we use cpu offload for training
fastvideo_args.use_cpu_offload = False
# make sure we are in training mode
fastvideo_args.inference_mode = False
# we hijack the precision to be the master weight type so that the
# model is loaded with the correct precision. Subsequently we will
# use FSDP2's MixedPrecisionPolicy to set the precision for the
# fwd, bwd, and other operations' precision.
fastvideo_args.precision = fastvideo_args.master_weight_type
assert fastvideo_args.precision == 'fp32', 'only fp32 is supported for training'
fastvideo_args.check_fastvideo_args()
logger.info(f"fastvideo_args in from_pretrained: {fastvideo_args}")
# fastvideo_args = FastVideoArgs(
# model_path=model_path,
# device_str=device or "cuda" if torch.cuda.is_available() else "cpu",
# **config_args)
fastvideo_args.check_fastvideo_args()
return cls(model_path,
fastvideo_args,
required_config_modules=required_config_modules)
@@ -140,7 +140,6 @@ class PreprocessPipeline(ComposedPipelineBase):
0], result_batch.prompt_attention_mask[0]
assert prompt_embeds.shape[0] == prompt_attention_mask.shape[0]
# Remove padding from prompt_embeds using attention mask for all batches
# Get sequence lengths from attention masks (number of 1s)
seq_lens = prompt_attention_mask.sum(dim=1)
# Create a list to store non-padded embeddings and masks
@@ -354,9 +353,6 @@ class PreprocessPipeline(ComposedPipelineBase):
"validation_parquet_dataset")
os.makedirs(validation_parquet_dir, exist_ok=True)
# Initialize Parquet dataset
validation_parquet_path = os.path.join(validation_parquet_dir,
"data.parquet")
with open(args.validation_prompt_txt, encoding="utf-8") as file:
lines = file.readlines()
+2 -3
View File
@@ -74,8 +74,7 @@ class DenoisingStage(PipelineStage):
)
# Setup precision and autocast settings
# target_dtype = PRECISION_TO_TYPE[fastvideo_args.precision]
target_dtype = torch.bfloat16
target_dtype = PRECISION_TO_TYPE[fastvideo_args.precision]
autocast_enabled = (target_dtype != torch.float32
) and not fastvideo_args.disable_autocast
@@ -190,7 +189,7 @@ class DenoisingStage(PipelineStage):
# Predict noise residual
with torch.autocast(device_type="cuda",
dtype=torch.bfloat16,
dtype=target_dtype,
enabled=autocast_enabled):
# TODO(will-refactor): all of this should be in the stage's init
+7 -11
View File
@@ -191,15 +191,14 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
logger.info(f"infos: {infos}")
caption = infos['caption']
captions.append(caption)
prompt_embeds = embeddings.to(fastvideo_args.device).to(torch.bfloat16)
prompt_attention_mask = masks.to(fastvideo_args.device).to(torch.bfloat16)
prompt_embeds = embeddings.to(fastvideo_args.device)
prompt_attention_mask = masks.to(fastvideo_args.device)
# Calculate sizes
latents_size = [(sampling_param.num_frames - 1) // 4 + 1,
sampling_param.height // 8,
sampling_param.width // 8]
n_tokens = latents_size[0] * latents_size[1] * latents_size[2]
logger.info('embed dtype', prompt_embeds.dtype)
# Prepare batch for validation
# print('shape of embeddings', prompt_embeds.shape)
@@ -216,7 +215,7 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
width=args.num_width,
num_frames=args.num_frames,
# num_inference_steps=fastvideo_args.validation_sampling_steps,
num_inference_steps=50,
num_inference_steps=10,
# guidance_scale=fastvideo_args.validation_guidance_scale,
guidance_scale=1,
n_tokens=n_tokens,
@@ -226,11 +225,10 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
)
# Run validation inference
with torch.autocast("cuda", dtype=torch.bfloat16):
with torch.inference_mode():
output_batch = self.validation_pipeline.forward(
batch, fastvideo_args)
samples = output_batch.output
with torch.inference_mode():
output_batch = self.validation_pipeline.forward(
batch, fastvideo_args)
samples = output_batch.output
# Process outputs
video = rearrange(samples, "b c t h w -> t b c h w")
@@ -531,8 +529,6 @@ class WanTrainingPipeline(TrainingPipeline):
args_copy.inference_mode = True
args_copy.vae_config.load_encoder = False
# TODO(will): clean this up
args_copy.precision = "bf16"
validation_pipeline = WanValidationPipeline.from_pretrained(
args.model_path, args=args_copy)
@@ -15,6 +15,7 @@ from fastvideo.v1.pipelines.stages import (ConditioningStage, DecodingStage,
TextEncodingStage,
TimestepPreparationStage)
# TODO(will): move PRECISION_TO_TYPE to better place
logger = init_logger(__name__)
+2 -2
View File
@@ -49,5 +49,5 @@ torchrun --nnodes 1 --nproc_per_node $NUM_GPUS\
--multi_phased_distill_schedule "4000-1" \
--weight_decay 0.01 \
--not_apply_cfg_solver \
--master_weight_type "fp32" \
--max_grad_norm 1.0
--master_weight_type "bf16" \
--max_grad_norm 1.0