sad / configs /sad_owt_b32_top1.yaml
haochengsama's picture
Add files using upload-large-folder tool
8b0aeb2 verified
Raw
History Blame
1.43 kB
# Block-wise AR Diffusion Config
#
# Key characteristics:
# - 序列分块:512 tokens / 8 per block = 64 blocks
# - 块内:并行扩散(双向 attention)
# - 块间:自回归(只能看到前面的块)
# - 训练:直接使用全 512 序列(不做课程学习)
model:
vocab_size: 50257
hidden_size: 768
n_blocks: 12
n_heads: 12
cond_dim: 128
max_seq_len: 512
block_size: 32
dropout: 0.0
num_levels: 2 # states: [V, 128]; mask = implicit level 2
level_sizes: [50257, 128]
ancestor:
lut_path: data/ancestor_lut_50257-128_top1_t1.0.pt
proto_path: data/hierarchy_prototypes_50257-128.pt
loss:
lambda_ancestor: 0.0 # set > 0 to enable ancestor CE loss
mask_only: true
training:
seed: 0
batch_size: 64
num_steps: 1_000_000
lr: 3.0e-4
lr_min: 3.0e-5
warmup_steps: 2000
weight_decay: 0.01
grad_clip: 1.0
dtype: bf16
compile: default # set to "off" to disable torch.compile
log_interval: 100
eval_interval: 5000
save_interval: 10000
data:
dataset: openwebtext
seq_len: 512
cache_dir: data/owt_cache
num_workers: 4
max_train_samples: null
max_val_samples: 100000
mode: subsample # subsample: 1 doc/sample, BOS/EOS, random window + pad (HDLM 对齐)
# pack: 跨文档拼接切块(旧行为)
logging:
use_wandb: true
project: sad_b32_top1
save_dir: outputs/sad_b32_top1