Potpov commited on
Commit
cebb108
·
verified ·
1 Parent(s): 80a071c

Upload mnist_color/art/config.yaml with huggingface_hub

Browse files
Files changed (1) hide show
  1. 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)