BlockDiffuse / config.yaml
tahamajs's picture
Upload configuration
c750327 verified
Raw
History Blame Contribute Delete
2.76 kB
# ==============================================================================
# BlockDiffuse Improved GPU Profile (Transfer Learning, CE, NN Loss, Curriculum)
# ==============================================================================
# Base LLM Metadata
base_llm:
model_name_or_path: "Qwen/Qwen2.5-0.5B-Instruct"
mid_layer_idx: 12
d_model: 896
num_layers: 24
num_heads: 14
vocab_size: 151936
max_prompt_len: 512
max_target_len: 100
# High-Capacity DiT Architecture with Transfer Learning & Deep Projection Head
dit:
d_model: 896
num_layers: 10 # Deeper 8-layer DiT model
num_heads: 14 # 14 attention heads (head_dim = 64)
mlp_ratio: 8.0
block_length: 100 # Maximum 100-token parallel generation
max_seq_len: 1024
dropout: 0.1
use_rope: true
use_block_causal: false
gradient_checkpointing: true # Dramatically cuts activation memory (~2GB VRAM)
adaln_zero: true
time_embedding_dim: 256
apply_rmsnorm_head: true
use_projection_head: true
projection_head_depth: "deep" # 3-layer deep MLP with 4x hidden expansion
init_from_base: true # Layer initialization from Qwen mid-layers (6-11)
init_num_layers: 6
init_start_layer: 6
init_base_model_name: "Qwen/Qwen2.5-0.5B-Instruct"
# Training Hyperparameters with Block-Length Curriculum
training:
learning_rate: 7.0e-4
min_lr: 1.0e-6
weight_decay: 0.01
warmup_steps: 200
max_steps: 5000
batch_size: 16 # 4x reduced batch size (effective batch = 8 with grad accum 2)
gradient_accumulation_steps: 1 # Effective global batch size = 32
max_grad_norm: 1.0
precision: "bfloat16" # High performance native bfloat16
save_interval: 1000
eval_interval: 200
log_interval: 10
output_dir: "./checkpoints_improved"
# Gradual Block-Length Curriculum
block_length_curriculum: true
curriculum_start_steps: 0
curriculum_increase_interval: 2000
curriculum_factor: 2
curriculum_min_block: 20
# Multi-Objective Loss with Discrete Token CE & Contrastive NN Supervision
loss_weights:
lambda_fm: 1.0 # Rectified flow matching velocity MSE (L_FM)
lambda_disp: 0.1 # Dispersive variance regularizer (L_Disp)
lambda_kl: 0.1 # Teacher Logit Distillation Loss (L_KL)
use_kldiv_loss: true
kl_temperature: 2.0
lambda_ce: 1.0 # Discrete Cross-Entropy Loss on Tokens (L_CE)
use_ce_loss: true
lambda_nn: 0.1 # Contrastive InfoNCE Nearest-Neighbor Loss (L_NN)
use_nn_loss: true
nn_temperature: 0.1
flow_matching:
sigma_min: 1.0e-5
time_sampling: "uniform"