Ouzhang's picture
Add files using upload-large-folder tool
3cd1076 verified
Raw
History Blame Contribute Delete
3.38 kB
defaults:
- _self_
- /callbacks: [checkpoint_every_n_steps, checkpoint_monitor, learning_rate_monitor]
- /data: openwebtext # switch to synthetic_alpha8 for debugging: set to /data: synthetic_alpha8
- /model: small
- /strategy: ddp
- /noise: log-linear
- /lr_scheduler: constant_warmup
- /prior: none
- /algo: flm
mode: train # train / ppl_eval / sample_eval
# di4c configs:
is_di4c: False
is_di4c_deterministic: True
seed: 1
loader:
global_batch_size: 512
eval_global_batch_size: ${.global_batch_size}
# Note: batch_size and eval_batch_size are **per machine**
batch_size: ${div_up:${.global_batch_size}, ${eval:${trainer.devices} * ${trainer.num_nodes}}}
eval_batch_size: ${div_up:${.eval_global_batch_size}, ${eval:${trainer.devices} * ${trainer.num_nodes}}}
num_workers: 4
pin_memory: True
sampling:
predictor: ancestral # ancestral_cache (only for MDLM), ancestral, analytic
steps: [1024]
noise_removal: ancestral # 'ancestral', 'greedy', 'none', 'meanflow'
use_float64: True
p_nucleus: 1.0
num_sample_batches: 1 # Total samples: `num_gpus` * `loader.eval_batch_size` * num_sample_batches
num_sample_log: 2
semi_ar: False
stride_length: 1
num_strides: 1
num_reflow_samples: 50000 # for generating reflow dataset
temperature: 1.0
solver: euler
gamma: 0.0
training:
loss_type: flow
pred_type: x0
ema: 0.9999
antithetic_sampling: True
importance_sampling: False
sampling_eps: 1e-3
change_of_variables: False
loss_precision: 'bf16' # bf16, float32, float64
finetune_path: ''
not_sampling_t: False
load_ema: False # load EMA parameters when finetuning
t_curriculum_steps: 100000
t_curriculum_min: 0.00
t_curriculum_max: 1.00
eval:
checkpoint_path: '' # Used to evaluate a checkpoint after training.
disable_ema: False
ema_decay: null # If set, select matching EMA from checkpoint
compute_generative_perplexity: True
perplexity_batch_size: 8
compute_perplexity_on_sanity: False
gen_ppl_eval_model_name_or_path: gpt2-large # gpt2-large, meta-llama/Llama-2-7b-hf, meta-llama/Llama-3.1-8B
generate_samples: True
generated_samples_path: ${cwd:}/samples.json
optim:
weight_decay: 0
lr: 3e-4
beta1: 0.9
beta2: 0.999
eps: 1e-8
ln_tune: none
trainer:
_target_: lightning.Trainer
accelerator: cuda
num_nodes: 1
devices: ${device_count:}
accumulate_grad_batches: ${div_up:${loader.global_batch_size}, ${eval:${trainer.devices} * ${loader.batch_size} * ${trainer.num_nodes}}}
gradient_clip_val: 1.0
precision: 'bf16'
num_sanity_val_steps: 2
max_steps: 1_000_000
log_every_n_steps: 100
limit_train_batches: 1.0 # train on full dataset, can be used to toggle quick run
limit_val_batches: 1.0 # validate on full dataset, can be used to toggle quick run
val_check_interval: 5000
wandb:
project: flm
notes: FLM
group: null
job_type: null
name: null
id: ${.name}_${seed}
tags:
- ${data.train}
- ${algo.name}
hydra:
run:
dir: ./outputs/${data.train}/${now:%Y.%m.%d}/${now:%H%M%S}
job:
chdir: true
checkpointing:
# Use custom `save_dir` if, e.g., saving to S3 bucket, otherwise leave this parameter as is
save_dir: ${cwd:}
# Note: `checkpoints` path should correspond to `checkpoint_every_n_steps.dirpath`
resume_from_ckpt: false
resume_ckpt_path: ${.save_dir}/checkpoints/last.ckpt