{"id":19961300,"url":"https://github.com/damo-nlp-sg/digit","last_synced_at":"2025-05-03T22:30:42.289Z","repository":{"id":258441546,"uuid":"872760851","full_name":"DAMO-NLP-SG/DiGIT","owner":"DAMO-NLP-SG","description":"[NeurIPS 2024] Stabilize the Latent Space for Image Autoregressive Modeling: A Unified Perspective","archived":false,"fork":false,"pushed_at":"2024-10-21T11:39:40.000Z","size":15200,"stargazers_count":30,"open_issues_count":0,"forks_count":2,"subscribers_count":5,"default_branch":"main","last_synced_at":"2024-10-23T03:38:38.282Z","etag":null,"topics":["autoregressive","fairseq","gpt","image-generation","language-model","neurips","transformer"],"latest_commit_sha":null,"homepage":"","language":"Python","has_issues":true,"has_wiki":null,"has_pages":null,"mirror_url":null,"source_name":null,"license":"mit","status":null,"scm":"git","pull_requests_enabled":true,"icon_url":"https://github.com/DAMO-NLP-SG.png","metadata":{"files":{"readme":"README.md","changelog":null,"contributing":null,"funding":null,"license":"LICENSE","code_of_conduct":null,"threat_model":null,"audit":null,"citation":null,"codeowners":null,"security":null,"support":null,"governance":null,"roadmap":null,"authors":null,"dei":null,"publiccode":null,"codemeta":null}},"created_at":"2024-10-15T03:02:22.000Z","updated_at":"2024-10-22T19:00:08.000Z","dependencies_parsed_at":null,"dependency_job_id":"3ff367de-f98c-4348-a2e1-65a103bb0003","html_url":"https://github.com/DAMO-NLP-SG/DiGIT","commit_stats":null,"previous_names":["damo-nlp-sg/digit"],"tags_count":0,"template":false,"template_full_name":null,"repository_url":"https://repos.ecosyste.ms/api/v1/hosts/GitHub/repositories/DAMO-NLP-SG%2FDiGIT","tags_url":"https://repos.ecosyste.ms/api/v1/hosts/GitHub/repositories/DAMO-NLP-SG%2FDiGIT/tags","releases_url":"https://repos.ecosyste.ms/api/v1/hosts/GitHub/repositories/DAMO-NLP-SG%2FDiGIT/releases","manifests_url":"https://repos.ecosyste.ms/api/v1/hosts/GitHub/repositories/DAMO-NLP-SG%2FDiGIT/manifests","owner_url":"https://repos.ecosyste.ms/api/v1/hosts/GitHub/owners/DAMO-NLP-SG","download_url":"https://codeload.github.com/DAMO-NLP-SG/DiGIT/tar.gz/refs/heads/main","host":{"name":"GitHub","url":"https://github.com","kind":"github","repositories_count":224374779,"owners_count":17300691,"icon_url":"https://github.com/github.png","version":null,"created_at":"2022-05-30T11:31:42.601Z","updated_at":"2022-07-04T15:15:14.044Z","host_url":"https://repos.ecosyste.ms/api/v1/hosts/GitHub","repositories_url":"https://repos.ecosyste.ms/api/v1/hosts/GitHub/repositories","repository_names_url":"https://repos.ecosyste.ms/api/v1/hosts/GitHub/repository_names","owners_url":"https://repos.ecosyste.ms/api/v1/hosts/GitHub/owners"}},"keywords":["autoregressive","fairseq","gpt","image-generation","language-model","neurips","transformer"],"created_at":"2024-11-13T02:07:10.771Z","updated_at":"2024-11-13T02:07:11.376Z","avatar_url":"https://github.com/DAMO-NLP-SG.png","language":"Python","funding_links":[],"categories":[],"sub_categories":[],"readme":"\u003ch1 align=\"center\"\u003e Stabilize the Latent Space for Image Autoregressive Modeling: A Unified Perspective (NeurIPS 2024)\n\u003c/h1\u003e\n\n\n\u003cdiv align=\"center\"\u003e\n\n[![arXiv](https://img.shields.io/badge/arXiv%20paper-2410.12490-b31b1b.svg)](https://arxiv.org/abs/2410.12490)\u0026nbsp;\n[![benchmark](https://img.shields.io/badge/Rank%204-Image%20Generation%20on%20ImageNet%20%28AR%29-32B1B4?logo=data%3Aimage%2Fsvg%2Bxml%3Bbase64%2CPHN2ZyB3aWR0aD0iNjA2IiBoZWlnaHQ9IjYwNiIgeG1sbnM9Imh0dHA6Ly93d3cudzMub3JnLzIwMDAvc3ZnIiB4bWxuczp4bGluaz0iaHR0cDovL3d3dy53My5vcmcvMTk5OS94bGluayIgb3ZlcmZsb3c9ImhpZGRlbiI%2BPGRlZnM%2BPGNsaXBQYXRoIGlkPSJjbGlwMCI%2BPHJlY3QgeD0iLTEiIHk9Ii0xIiB3aWR0aD0iNjA2IiBoZWlnaHQ9IjYwNiIvPjwvY2xpcFBhdGg%2BPC9kZWZzPjxnIGNsaXAtcGF0aD0idXJsKCNjbGlwMCkiIHRyYW5zZm9ybT0idHJhbnNsYXRlKDEgMSkiPjxyZWN0IHg9IjUyOSIgeT0iNjYiIHdpZHRoPSI1NiIgaGVpZ2h0PSI0NzMiIGZpbGw9IiM0NEYyRjYiLz48cmVjdCB4PSIxOSIgeT0iNjYiIHdpZHRoPSI1NyIgaGVpZ2h0PSI0NzMiIGZpbGw9IiM0NEYyRjYiLz48cmVjdCB4PSIyNzQiIHk9IjE1MSIgd2lkdGg9IjU3IiBoZWlnaHQ9IjMwMiIgZmlsbD0iIzQ0RjJGNiIvPjxyZWN0IHg9IjEwNCIgeT0iMTUxIiB3aWR0aD0iNTciIGhlaWdodD0iMzAyIiBmaWxsPSIjNDRGMkY2Ii8%2BPHJlY3QgeD0iNDQ0IiB5PSIxNTEiIHdpZHRoPSI1NyIgaGVpZ2h0PSIzMDIiIGZpbGw9IiM0NEYyRjYiLz48cmVjdCB4PSIzNTkiIHk9IjE3MCIgd2lkdGg9IjU2IiBoZWlnaHQ9IjI2NCIgZmlsbD0iIzQ0RjJGNiIvPjxyZWN0IHg9IjE4OCIgeT0iMTcwIiB3aWR0aD0iNTciIGhlaWdodD0iMjY0IiBmaWxsPSIjNDRGMkY2Ii8%2BPHJlY3QgeD0iNzYiIHk9IjY2IiB3aWR0aD0iNDciIGhlaWdodD0iNTciIGZpbGw9IiM0NEYyRjYiLz48cmVjdCB4PSI0ODIiIHk9IjY2IiB3aWR0aD0iNDciIGhlaWdodD0iNTciIGZpbGw9IiM0NEYyRjYiLz48cmVjdCB4PSI3NiIgeT0iNDgyIiB3aWR0aD0iNDciIGhlaWdodD0iNTciIGZpbGw9IiM0NEYyRjYiLz48cmVjdCB4PSI0ODIiIHk9IjQ4MiIgd2lkdGg9IjQ3IiBoZWlnaHQ9IjU3IiBmaWxsPSIjNDRGMkY2Ii8%2BPC9nPjwvc3ZnPg%3D%3D)](https://paperswithcode.com/sota/image-generation-on-imagenet-256x256?tag_filter=485\u0026p=stabilize-the-latent-space-for-image)\n\n\u003c/div\u003e\n\n\n![FID_IS](assets/FID_IS.png)\n\n\n## Overview\n\n![The overview of DiGIT](assets/digit_model.png)\n\nWe present **DiGIT**, an auto-regressive generative model performing next-token prediction in an abstract latent space derived from self-supervised learning (SSL) models. By employing K-Means clustering on the hidden states of the DINOv2 model, we effectively create a novel discrete tokenizer. This method significantly boosts image generation performance on ImageNet dataset, achieving an FID score of **4.59 for class-unconditional tasks** and **3.39 for class-conditional tasks**. Additionally, the model enhances image understanding, achieving a **linear-probe accuracy of 80.3**.\n\n\n## Experimental Results\n\n### Linear-Probe Accuracy on ImageNet\n\n\n| Methods                          | \\# Tokens   | Features | \\# Params  | Top-1 Acc. $\\uparrow$ |\n|-----------------------------------|-------------|----------|------------|-----------------------|\n| iGPT-L   | 32 $\\times$ 32 | 1536     | 1362M      | 60.3                  |\n| iGPT-XL   | 64 $\\times$ 64 | 3072     | 6801M      | 68.7                  |\n| VIM+VQGAN  | 32 $\\times$ 32 | 1024     | 650M       | 61.8                  |\n| VIM+dVAE  | 32 $\\times$ 32 | 1024     | 650M       | 63.8                  |\n| VIM+ViT-VQGAN  | 32 $\\times$ 32 | 1024     | 650M       | 65.1                  |\n| VIM+ViT-VQGAN  | 32 $\\times$ 32 | 2048     | 1697M      | 73.2                  |\n| AIM          | 16 $\\times$ 16 | 1536     | 0.6B       | 70.5                  |\n| **DiGIT (Ours)**                  | 16 $\\times$ 16 | 1024     | 219M       | 71.7                  |\n| **DiGIT (Ours)**                  | 16 $\\times$ 16 | 1536     | 732M       | **80.3**               |\n\n### Class-Unconditional Image Generation on ImageNet (Resolution: 256 $\\times$ 256)\n\n| Type  | Methods                             | \\# Param | \\# Epoch | FID $\\downarrow$ | IS $\\uparrow$  |\n|-------|-------------------------------------|----------|----------|------------------|----------------|\n| GAN   | BigGAN        | 70M      | -        | 38.6             | 24.70          |\n| Diff. | LDM          | 395M     | -        | 39.1             | 22.83          |\n| Diff. | ADM     | 554M     | -        | 26.2             | 39.70          |\n| MIM   | MAGE              | 200M     | 1600     | 11.1             | 81.17          |\n| MIM   | MAGE              | 463M     | 1600     | 9.10             | 105.1          |\n| MIM   | MaskGIT      | 227M     | 300      | 20.7             | 42.08          |\n| MIM   | **DiGIT (+MaskGIT)**                | 219M     | 200      | **9.04**         | **75.04**      |\n| AR    | VQGAN         | 214M     | 200      | 24.38            | 30.93          |\n| AR    | **DiGIT (+VQGAN)**                  | 219M     | 400      | **9.13**         | **73.85**      |\n| AR    | **DiGIT (+VQGAN)**                  | 732M     | 200      | **4.59**         | **141.29**     |\n\n### Class-Conditional Image Generation on ImageNet (Resolution: 256 $\\times$ 256)\n\n\n\n| Type  | Methods              | \\# Param | \\# Epoch | FID $\\downarrow$ | IS $\\uparrow$  |\n|-------|----------------------|----------|----------|------------------|----------------|\n| GAN   | BigGAN               | 160M     | -        | 6.95             | 198.2          |\n| Diff. | ADM                  | 554M     | -        | 10.94            | 101.0          |\n| Diff. | LDM-4                | 400M     | -        | 10.56            | 103.5          |\n| Diff. | DiT-XL/2             | 675M     | -        | 9.62             | 121.50         |\n| Diff. | L-DiT-7B             | 7B       | -        | 6.09             | 153.32         |\n| MIM   | CQR-Trans            | 371M     | 300      | 5.45             | 172.6          |\n| MIM+AR | VAR                 | 310M     | 200      | 4.64             | -              |\n| MIM+AR | VAR                 | 310M     | 200      | 3.60* | 257.5* |\n| MIM+AR | VAR                 | 600M     | 250      | 2.95* | 306.1* |\n| MIM   | MAGVIT-v2            | 307M     | 1080     | 3.65             | 200.5          |\n| AR    | VQVAE-2              | 13.5B    | -        | 31.11            | 45             |\n| AR    | RQ-Trans             | 480M     | -        | 15.72            | 86.8           |\n| AR    | RQ-Trans             | 3.8B     | -        | 7.55             | 134.0          |\n| AR    | ViTVQGAN             | 650M     | 360      | 11.20            | 97.2           |\n| AR    | ViTVQGAN             | 1.7B     | 360      | 5.3              | 149.9          |\n| MIM   | MaskGIT              | 227M     | 300      | 6.18             | 182.1          |\n| MIM   | **DiGIT (+MaskGIT)** | 219M     | 200      | **4.62**         | **146.19**     |\n| AR    | VQGAN                | 227M     | 300      | 18.65            | 80.4           |\n| AR    | **DiGIT (+VQGAN)**   | 219M     | 400      | **4.79**         | **142.87**     |\n| AR    | **DiGIT (+VQGAN)**   | 732M     | 200      | **3.39**         | **205.96**     |\n\n*: VAR is trained with classifier-free guidance while all the other models are not.\n\n\n## Checkpoints\nThe K-Means npy file and model checkpoints can be downloaded from: \n\n|   Model    | Link |            \n|:----------:|:-----:|\n|  HF weights🤗    |  [Huggingface](https://huggingface.co/DAMO-NLP-SG/DiGIT) |\n\n\nFor the base model we use [DINOv2-base](https://dl.fbaipublicfiles.com/dinov2/dinov2_vitb14/dinov2_vitb14_reg4_pretrain.pth) and [DINOv2-large](https://dl.fbaipublicfiles.com/dinov2/dinov2_vitl14/dinov2_vitl14_reg4_pretrain.pth) for large size model. The VQGAN we use is the same as [MAGE](https://drive.google.com/file/d/13S_unB87n6KKuuMdyMnyExW0G1kplTbP/view?usp=sharing).\n\n\n\n```\nDiGIT\n└── data/\n    ├── ILSVRC2012\n        ├── dinov2_base_short_224_l3\n            ├── km_8k.npy\n        ├── dinov2_large_short_224_l3\n            ├── km_16k.npy\n└── outputs/\n    ├── base_8k_stage1\n    ├── ...\n└── models/\n    ├── vqgan_jax_strongaug.ckpt\n    ├── dinov2_vitb14_reg4_pretrain.pth\n    ├── dinov2_vitl14_reg4_pretrain.pth\n```\n\n\n## Preparation\n\n### Installation\n1. Download the code\n```shell \ngit clone https://github.com/DAMO-NLP-SG/DiGIT.git\ncd DiGIT\n```\n\n2. Install `fairseq` via `pip install fairseq`.\n\n\n### Dataset Preparation\nDownload [ImageNet](http://image-net.org/) dataset, and place it in your dataset dir `$PATH_TO_YOUR_WORKSPACE/dataset/ILSVRC2012`. \n\n### Tokenizer\nExtract SSL features and save them as .npy files. Use the K-Means algorithm with [faiss](https://github.com/facebookresearch/faiss) to compute the centroids. You can also utilize our pre-trained centroids available on [Huggingface](https://huggingface.co/DAMO-NLP-SG/DiGIT).\n\n```shell\nbash preprocess/run.sh\n```\n\n### Training Scripts \n\n**Step1**\n\nTrain a GPT model with a discriminative tokenizer. You can find the training scripts in `scripts/train_stage1_ar.sh` and the hyper-params are in `config/stage1/dino_base.yaml`. For class conditional generation configuration, see `scripts/train_stage1_classcond.sh`.\n\n**Step2**\n\nTrain a pixel decoder (either AR model or NAR model) conditioned on the discriminative tokens. You can find the autoregressive training scripts in `scripts/train_stage2_ar.sh` and NAR training scripts in `scripts/train_stage2_nar.sh`.\n\nA folder named `outputs/EXP_NAME/checkpoints` will be created to save the checkpoints. TensorBoard log files are saved at `outputs/EXP_NAME/tb`. Logs will be recorded in `outputs/EXP_NAME/train.log`. \n\nYou can monitor the training process using `tensorboard --logdir=outputs/EXP_NAME/tb`.\n\n\n### Sampling Scripts\n\nFirst sampling discriminative tokens with `scripts/infer_stage1_ar.sh`. For the base model size, we recommend setting topk=200, and for a large model size, use topk=400.\n\nThen run `scripts/infer_stage2_ar.sh` to sample VQ tokens based on the previously sampled discriminative tokens.\n\nGenerated tokens and synthesized images will be stored in a directory named `outputs/EXP_NAME/results`.\n\n### FID and IS evaluation\nPrepare the ImageNet validation set for FID evaluation:\n```shell\npython prepare_imgnet_val.py --data_path $PATH_TO_YOUR_WORKSPACE/dataset/ILSVRC2012 --output_dir imagenet-val\n```\n\nInstall the evaluation tool by running `pip install torch-fidelity`.\n\nExecute the following command to evaluate FID:\n```shell \npython fairseq_user/eval_fid.py --results-path $IMG_SAVE_DIR --subset $GEN_SUBSET\n```\n\n### Linear Probe training\n\n```shell\nbash scripts/train_stage1_linearprobe.sh\n```\n\n## License\nThis project is licensed under the MIT License - see the [LICENSE](LICENSE) file for details.\n\n\n## Citation\n\nIf you find our project useful, hope you can star our repo and cite our work as follows. \n\n```bibtex\n@misc{zhu2024stabilize,\n    title={Stabilize the Latent Space for Image Autoregressive Modeling: A Unified Perspective},\n    author={Yongxin Zhu and Bocheng Li and Hang Zhang and Xin Li and Linli Xu and Lidong Bing},\n    year={2024},\n    eprint={2410.12490},\n    archivePrefix={arXiv},\n    primaryClass={cs.CV}\n}\n```\n","project_url":"https://awesome.ecosyste.ms/api/v1/projects/github.com%2Fdamo-nlp-sg%2Fdigit","html_url":"https://awesome.ecosyste.ms/projects/github.com%2Fdamo-nlp-sg%2Fdigit","lists_url":"https://awesome.ecosyste.ms/api/v1/projects/github.com%2Fdamo-nlp-sg%2Fdigit/lists"}