Upload mnist_color/art/config.yaml with huggingface_hub
Browse files- mnist_color/art/config.yaml +65 -0
mnist_color/art/config.yaml
ADDED
|
@@ -0,0 +1,65 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
stage1_params:
|
| 2 |
+
name: "VSQ" # always keep this "VectorGPT"
|
| 3 |
+
vector_decoder_model: "cnn" # "mlp" or "raster_conv"
|
| 4 |
+
quantized_dim: 512
|
| 5 |
+
codebook_size: 4096 # will be ignored for FSQ
|
| 6 |
+
image_loss: "pyramid" # "pyramid" or "mse"
|
| 7 |
+
single_code_representation: true
|
| 8 |
+
vq_method: "fsq" # "vqvae", "FSQ", "vqtorch"
|
| 9 |
+
fsq_levels: [7,5,5,5,5] # will determine codebook_size, see Table 1 of FSQ paper - [7,5,5,5,5] for 4096 (e.g. StrokeNUWA), [8,5,5,5] for 1024
|
| 10 |
+
num_segments: 15
|
| 11 |
+
pred_color: true
|
| 12 |
+
checkpoint_path: "/sc/projects/sci-aisc/marco.cipriano/results/svg/Grimoire/VSQ/TiledMNIST/checkpoints/last.ckpt"
|
| 13 |
+
|
| 14 |
+
|
| 15 |
+
model_params:
|
| 16 |
+
name: "VQ_SVG_Stage2"
|
| 17 |
+
max_seq_len: 512 # 512
|
| 18 |
+
dim: 512 # 512
|
| 19 |
+
depth: 16 # 12
|
| 20 |
+
heads: 8
|
| 21 |
+
text_encoder_str: "bert-base-uncased"
|
| 22 |
+
use_alibi_positional_bias : False
|
| 23 |
+
|
| 24 |
+
data_params:
|
| 25 |
+
dataset: "mnist"
|
| 26 |
+
csv_path: "/sc/projects/sci-aisc/marco.cipriano/data/SVG/Grimoire/MNIST/8x8_randomcolor/split.csv" # root path for CausalSVGDataModule in dataset.py
|
| 27 |
+
vq_token_npy_path: "/sc/projects/sci-aisc/marco.cipriano/data/SVG/Grimoire/MNIST/8x8_randomcolor/vsq_tokenized.npy"
|
| 28 |
+
train_batch_size: 16 #32
|
| 29 |
+
val_batch_size: 16 # 32
|
| 30 |
+
num_workers: 8
|
| 31 |
+
patch_size: 34
|
| 32 |
+
width: 128
|
| 33 |
+
num_tiles_per_row: 3
|
| 34 |
+
min_context_length: 10 # all samples below this will be removed from the dataset
|
| 35 |
+
fraction_of_strokenuwa_inputs: 0.0
|
| 36 |
+
fraction_of_class_only_inputs: 0.9 # fraction of samples that will only have the "class" entry of the dataframe as input
|
| 37 |
+
fraction_of_blank_inputs: 0.1 # fraction of samples that will have empty text input
|
| 38 |
+
fraction_of_iconshop_chatgpt_inputs: 0.0
|
| 39 |
+
shuffle_vq_order: False # whether to shuffle the order of the VQ tokens, its not really "shuffling", but more cutting the sequence into two parts and switching their order
|
| 40 |
+
use_pre_computed_text_tokens_only: False # if True, the text tokens that were pre-computed will be used as input, if False, the text will be tokenized in the dataloader (according to the specified fractions).
|
| 41 |
+
|
| 42 |
+
exp_params:
|
| 43 |
+
lr: 0.0006 # 0.00002
|
| 44 |
+
weight_decay: 1.e-4 # specify positive float to enable, start experimenting with 1.e-4/1.e-3
|
| 45 |
+
scheduler_gamma: 0.96 # 0.95 is a good starting value
|
| 46 |
+
train_log_interval: 0.05 # decides how often full generations are made and logged
|
| 47 |
+
val_log_interval: 0.02
|
| 48 |
+
metric_log_interval: 0.4
|
| 49 |
+
manual_seed: 1265
|
| 50 |
+
post_process: False # whether to log with svg fixing
|
| 51 |
+
|
| 52 |
+
trainer_params:
|
| 53 |
+
devices: -1 # always keep at -1 as this takes all available GPUs specified through CUDA_VISIBLE_DEVICES
|
| 54 |
+
max_epochs: 100 # dsnt matter too much, got early stopping implemented
|
| 55 |
+
# accumulate_grad_batches: 2
|
| 56 |
+
|
| 57 |
+
logging_params:
|
| 58 |
+
entity: "aiis-chair" # comment to use default wandb entity "mfeuer"
|
| 59 |
+
project: "grimoire-2" # your wandb project name
|
| 60 |
+
save_dir: "/sc/projects/sci-aisc/marco.cipriano/results/svg/Grimoire/ART"
|
| 61 |
+
name: "Stage 2 - MNIST color" # name of the run in wandb
|
| 62 |
+
version: 0
|
| 63 |
+
author: "Marco" # will be a tag in wandb
|
| 64 |
+
id: null # id of wandb run to continue
|
| 65 |
+
allow_val_change: False # allow changing values in this config w.r.t. the run that you're continuing (good for changing loss weightings mid-run)
|