CoolFace
Modelpublic

msj19/gated_deltaproduct

sourceHugging Faceupdated 8mo agoView on Hugging Face
0likes5downloads
train_node.sh127 linesDownload Raw Back to root
1args=$@2for arg in $args; do3    eval "$arg"4done5 6echo "model:            ${model:=fla-hub/gla-1.3B-100B}"7echo "tokenizer:        ${tokenizer:=/mnt/jfzn/msj/delta_net-1.3B-100B}"8echo "project:          ${project:=fla}"9echo "type:             ${type:=gla}"10echo "data:             ${data:=}"11echo "name:             ${name:=}"12echo "cache:            ${cache:=}"13echo "varlen:           ${varlen:=false}"14echo "seed:             ${seed:=42}"15echo "context:          ${context:=2048}"16echo "steps:            ${steps:=0}"17echo "save:             ${save:=2048}"18echo "limit:            ${limit:=1}"19echo "preprocessing:    ${preprocessing:=32}"20echo "workers:          ${workers:=32}"21echo "prefetch:         ${prefetch:=2}"22echo "logging:          ${logging:=32}"23echo "config:           ${config:=configs/deepspeed_multi.yaml}"24 25echo "lr:               ${lr:=3e-4}"26echo "scheduler:        ${scheduler:=cosine_with_min_lr}"27echo "epochs:           ${epochs:=1}"28echo "optim:            ${optim:=adamw_torch_fused}"29echo "decay:            ${decay:=0.01}"30echo "beta1:            ${beta1:=0.9}"31echo "beta2:            ${beta2:=0.95}"32echo "norm:             ${norm:=1.0}"33echo "batch:            ${batch:=32}"34echo "update:           ${update:=1}"35echo "warmup:           ${warmup:=512}"36echo "path:             ${path:=}"37echo "checkpoint:       ${checkpoint:=}"38echo "node:             ${node:=}"39echo "rank:             ${rank:=}"40echo "ip:               ${ip:=10.119.141.222}"41echo "port:             ${port:=}"42echo "nodes:            ${nodes:=1}"43echo "gpus:             ${gpus:=8}"44 45params="--model_name_or_path $model \46    --tokenizer $tokenizer \47    --use_fast_tokenizer \48    --do_train \49    --dataset $data \50    --context_length $context \51    --preprocessing_num_workers $preprocessing \52    --dataloader_num_workers $workers \53    --dataloader_prefetch_factor $prefetch \54    --output_dir $path \55    --overwrite_output_dir \56    --logging_steps $logging \57    --include_num_input_tokens_seen \58    --save_steps $save \59    --save_total_limit $limit \60    --learning_rate $lr \61    --lr_scheduler_type $scheduler \62    --warmup_steps $warmup \63    --optim $optim \64    --weight_decay $decay \65    --adam_beta1=$beta1 \66    --adam_beta2=$beta2 \67    --max_grad_norm $norm \68    --num_train_epochs $epochs \69    --per_device_train_batch_size $batch \70    --gradient_accumulation_steps $update \71    --seed $seed \72    --logging_steps $logging \73    --log_level info \74    --bf16"75 76if [ $steps -gt 0 ]; then77    params+=" --max_steps $steps"78fi79 80if [ "$name" != "" ]; then81  params+=" --dataset_name $name"82fi83if [ "$cache" != "" ]; then84  params+=" --cache_dir $cache"85fi86if [ "$varlen" == "true" ]; then87  params+=" --varlen"88fi89if [ "$checkpoint" != "" ]; then90  params+=" --resume_from_checkpoint $checkpoint"91echo '*****************************************'$checkpoint92fi93# if [ "$WANDB_DISABLED" != "true" ]; then94#   params+=" --report_to wandb \95#   --run_name $type.$(basename $path)"96# else97params+=" --report_to none"98# fi99 100echo "Launching training..."101accelerate_params=""102if [ "$rank" != "" ]; then103  accelerate_params+=" --machine_rank $rank  \104    --num_processes $((nodes * gpus)) \105    --num_machines $nodes \106    --main_process_ip $ip \107    --main_process_port $port \108    --same_network"109fi110 111 112set -x113mkdir -p $path114cp * $path115cp -r configs $path116cp -r flame   $path117cp -r fla2 $path118cp -r fla3 $path119# export WANDB_DISABLED=1120export TRANSFORMERS_OFFLINE=1121export HF_DATASETS_OFFLINE=1122if [ "$date" == "" ]; then123  date=$(date +%Y%m%d%H%M)124fi125 126/mnt/jfzn/miniconda3/envs/msj_eval/bin/accelerate launch --config_file $config  run.py $params127echo "RUNNING DONE!"