msj19/gated_deltaproduct
05
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!"