11 KiB
SigLIPsiglip
๊ฐ์overview
SigLIP ๋ชจ๋ธ์ Xiaohua Zhai, Basil Mustafa, Alexander Kolesnikov, Lucas Beyer์ Sigmoid Loss for Language Image Pre-Training ๋ ผ๋ฌธ์์ ์ ์๋์์ต๋๋ค. SigLIP์ CLIP์์ ์ฌ์ฉ๋ ์์ค ํจ์๋ฅผ ๊ฐ๋จํ ์๋ณ ์๊ทธ๋ชจ์ด๋ ์์ค(pairwise sigmoid loss)๋ก ๋์ฒดํ ๊ฒ์ ์ ์ํฉ๋๋ค. ์ด๋ ImageNet์์ ์ ๋ก์ท ๋ถ๋ฅ ์ ํ๋ ์ธก๋ฉด์์ ๋ ๋์ ์ฑ๋ฅ์ ๋ณด์ ๋๋ค.
๋ ผ๋ฌธ์ ์ด๋ก์ ๋ค์๊ณผ ๊ฐ์ต๋๋ค:
์ฐ๋ฆฌ๋ ์ธ์ด-์ด๋ฏธ์ง ์ฌ์ ํ์ต(Language-Image Pre-training, SigLIP)์ ์ํ ๊ฐ๋จํ ์๋ณ ์๊ทธ๋ชจ์ด๋ ์์ค์ ์ ์ํฉ๋๋ค. ์ํํธ๋งฅ์ค ์ ๊ทํ๋ฅผ ์ฌ์ฉํ๋ ํ์ค ๋์กฐ ํ์ต๊ณผ ๋ฌ๋ฆฌ, ์๊ทธ๋ชจ์ด๋ ์์ค์ ์ด๋ฏธ์ง-ํ ์คํธ ์์๋ง ์์ฉํ๋ฉฐ ์ ๊ทํ๋ฅผ ์ํด ์๋ณ ์ ์ฌ์ฑ์ ์ ์ญ์ ๊ด์ ์ ํ์๋ก ํ์ง ์์ต๋๋ค. ์๊ทธ๋ชจ์ด๋ ์์ค์ ๋ฐฐ์น ํฌ๊ธฐ๋ฅผ ๋์ฑ ํ์ฅํ ์ ์๊ฒ ํ๋ ๋์์ ์์ ๋ฐฐ์น ํฌ๊ธฐ์์๋ ๋ ๋์ ์ฑ๋ฅ์ ๋ณด์ ๋๋ค. Locked-image Tuning๊ณผ ๊ฒฐํฉํ์ฌ, ๋จ 4๊ฐ์ TPUv4 ์นฉ๋ง์ผ๋ก ์ดํ ๋ง์ 84.5%์ ImageNet ์ ๋ก์ท ์ ํ๋๋ฅผ ๋ฌ์ฑํ๋ SigLiT ๋ชจ๋ธ์ ํ์ตํ์ต๋๋ค. ์์ค ํจ์์์ ๋ฐฐ์น ํฌ๊ธฐ๋ฅผ ๋ถ๋ฆฌํจ์ผ๋ก์จ ์์ ๋ ์์ ์ํฅ๊ณผ Negative ๋ Positive ๋น์จ์ ์ฐ๊ตฌํ ์ ์๊ฒ ๋์์ต๋๋ค. ๋ง์ง๋ง์ผ๋ก, ์ฐ๋ฆฌ๋ ๋ฐฐ์น ํฌ๊ธฐ๋ฅผ 100๋ง ๊ฐ๊น์ง ๊ทน๋จ์ ์ผ๋ก ๋๋ ค๋ณด์๊ณ , ๋ฐฐ์น ํฌ๊ธฐ ์ฆ๊ฐ์ ์ด์ ์ด ๋น ๋ฅด๊ฒ ๊ฐ์ํ๋ฉฐ 32k์ ๋ ํฉ๋ฆฌ์ ์ธ ๋ฐฐ์น ํฌ๊ธฐ๋ก๋ ์ถฉ๋ถํ๋ค๋ ๊ฒ์ ๋ฐ๊ฒฌํ์ต๋๋ค.
์ฌ์ฉ ํusage-tips
- SigLIP์ ์ฌ์ฉ๋ฒ์ CLIP๊ณผ ์ ์ฌํฉ๋๋ค. ์ฃผ์ ์ฐจ์ด์ ์ ํ์ต ์์ค ํจ์๋ก, ๋ฐฐ์น ๋ด ๋ชจ๋ ์ด๋ฏธ์ง์ ํ ์คํธ ๊ฐ์ ์๋ณ ์ ์ฌ์ฑ์ ๋ํ ์ ์ญ์ ๊ด์ ์ด ํ์ํ์ง ์์ต๋๋ค. ์ํํธ๋งฅ์ค ๋์ ๋ก์ง์ ์๊ทธ๋ชจ์ด๋ ํ์ฑํ ํจ์๋ฅผ ์ ์ฉํด์ผ ํฉ๋๋ค.
- ํ์ต์ ์ง์๋์ง๋ง
torch.distributed์ ํธ๋ฆฌํฐ๋ฅผ ์ฌ์ฉํ์ง ์์ ๋ฐฐ์น ํฌ๊ธฐ์ ํ์ฅ์ฑ์ด ์ ํ๋ ์ ์์ต๋๋ค. ๊ทธ๋ฌ๋ ๋จ์ผ ๋ ธ๋ ๋ค์ค GPU ์ค์ ์์๋ DDP์ FDSP๊ฐ ์๋ํฉ๋๋ค. - ๋
๋ฆฝํ [
SiglipTokenizer] ๋๋ [SiglipProcessor]๋ฅผ ์ฌ์ฉํ ๋๋ ๋ชจ๋ธ์ด ๊ทธ๋ ๊ฒ ํ์ต๋์์ผ๋ฏ๋กpadding="max_length"๋ฅผ ์ ๋ฌํด์ผ ํฉ๋๋ค. - ํ์ดํ๋ผ์ธ๊ณผ ๋์ผํ ๊ฒฐ๊ณผ๋ฅผ ์ป์ผ๋ ค๋ฉด "This is a photo of {label}."์ ํ๋กฌํํธ ํ ํ๋ฆฟ์ ์ฌ์ฉํด์ผ ํฉ๋๋ค.
CLIP๊ณผ ๋น๊ตํ SigLIP ํ๊ฐ ๊ฒฐ๊ณผ. ์๋ณธ ๋ ผ๋ฌธ์์ ๋ฐ์ท.
์ด ๋ชจ๋ธ์ nielsr๊ฐ ๊ธฐ์ฌํ์ต๋๋ค. ์๋ณธ ์ฝ๋๋ ์ฌ๊ธฐ์์ ์ฐพ์ ์ ์์ต๋๋ค.
์ฌ์ฉ ์์usage-example
SigLIP์ ์ฌ์ฉํ๋ ๋ฐฉ๋ฒ์๋ ๋ ๊ฐ์ง ์ฃผ์ ๋ฐฉ๋ฒ์ด ์์ต๋๋ค: ๋ชจ๋ ๋ณต์ก์ฑ์ ์ถ์ํํ๋ ํ์ดํ๋ผ์ธ API๋ฅผ ์ฌ์ฉํ๊ฑฐ๋, ์ง์ SiglipModel ํด๋์ค๋ฅผ ์ฌ์ฉํ๋ ๋ฐฉ๋ฒ์
๋๋ค.
ํ์ดํ๋ผ์ธ APIpipeline-API
ํ์ดํ๋ผ์ธ์ ์ฌ์ฉํ๋ฉด ๋ช ์ค์ ์ฝ๋๋ก ๋ชจ๋ธ์ ์ฌ์ฉํ ์ ์์ต๋๋ค:
>>> from transformers import pipeline
>>> from PIL import Image
>>> import requests
>>> # ํ์ดํ๋ผ์ธ ๋ก๋
>>> image_classifier = pipeline(task="zero-shot-image-classification", model="google/siglip-base-patch16-224")
>>> # ์ด๋ฏธ์ง ๋ก๋
>>> url = 'http://images.cocodataset.org/val2017/000000039769.jpg'
>>> image = Image.open(requests.get(url, stream=True).raw)
>>> # ์ถ๋ก
>>> candidate_labels = ["2 cats", "a plane", "a remote"]
>>> outputs = image_classifier(image, candidate_labels=candidate_labels)
>>> outputs = [{"score": round(output["score"], 4), "label": output["label"] } for output in outputs]
>>> print(outputs)
[{'score': 0.1979, 'label': '2 cats'}, {'score': 0.0, 'label': 'a remote'}, {'score': 0.0, 'label': 'a plane'}]
์ง์ ๋ชจ๋ธ ์ฌ์ฉํ๊ธฐusing-the-model-yourself
์ ์ฒ๋ฆฌ์ ํ์ฒ๋ฆฌ๋ฅผ ์ง์ ์ํํ๋ ค๋ฉด ๋ค์๊ณผ ๊ฐ์ด ํ๋ฉด ๋ฉ๋๋ค:
>>> from PIL import Image
>>> import requests
>>> from transformers import AutoProcessor, AutoModel
>>> import torch
>>> model = AutoModel.from_pretrained("google/siglip-base-patch16-224")
>>> processor = AutoProcessor.from_pretrained("google/siglip-base-patch16-224")
>>> url = "http://images.cocodataset.org/val2017/000000039769.jpg"
>>> image = Image.open(requests.get(url, stream=True).raw)
>>> candidate_labels = ["2 cats", "2 dogs"]
# ํ์ดํ๋ผ์ธ ํ๋กฌํํธ ํ
ํ๋ฆฟ์ ๋ฐ๋ผ ๋์ผํ ๊ฒฐ๊ณผ๋ฅผ ์ป์ต๋๋ค
>>> texts = [f'This is a photo of {label}.' for label in candidate_labels]
# ์ค์: ๋ชจ๋ธ์ด ์ด๋ ๊ฒ ํ์ต๋์์ผ๋ฏ๋ก `padding=max_length`๋ฅผ ์ ๋ฌํฉ๋๋ค
>>> inputs = processor(text=texts, images=image, padding="max_length", return_tensors="pt")
>>> with torch.no_grad():
... outputs = model(**inputs)
>>> logits_per_image = outputs.logits_per_image
>>> probs = torch.sigmoid(logits_per_image) # ์๊ทธ๋ชจ์ด๋ ํ์ฑํ ํจ์๋ฅผ ์ ์ฉํ ํ๋ฅ ์
๋๋ค
>>> print(f"{probs[0][0]:.1%} that image 0 is '{candidate_labels[0]}'")
19.8% that image 0 is '2 cats'
๋ฆฌ์์คresources
SigLIP์ ์์ํ๋ ๋ฐ ๋์์ด ๋๋ ๊ณต์ Hugging Face ๋ฐ ์ปค๋ฎค๋ํฐ(๐๋ก ํ์) ๋ฆฌ์์ค ๋ชฉ๋ก์ ๋๋ค.
- ์ ๋ก์ท ์ด๋ฏธ์ง ๋ถ๋ฅ ์์ ๊ฐ์ด๋
- SigLIP์ ๋ํ ๋ฐ๋ชจ ๋ ธํธ๋ถ์ ์ฌ๊ธฐ์์ ์ฐพ์ ์ ์์ต๋๋ค. ๐
์ฌ๊ธฐ์ ํฌํจ๋ ๋ฆฌ์์ค๋ฅผ ์ ์ถํ๋ ๋ฐ ๊ด์ฌ์ด ์์ผ์๋ฉด Pull Request๋ฅผ ์ด์ด์ฃผ์๋ฉด ๊ฒํ ํ๊ฒ ์ต๋๋ค! ๋ฆฌ์์ค๋ ์ด์์ ์ผ๋ก ๊ธฐ์กด ๋ฆฌ์์ค๋ฅผ ๋ณต์ ํ๋ ๋์ ์๋ก์ด ๊ฒ์ ๋ณด์ฌ์ฃผ์ด์ผ ํฉ๋๋ค.
SigLIP๊ณผ Flash Attention 2 ๊ฒฐํฉํ๊ธฐcombining-siglip-with-flash-attention-2
๋จผ์ Flash Attention 2์ ์ต์ ๋ฒ์ ์ ์ค์นํด์ผ ํฉ๋๋ค.
pip install -U flash-attn --no-build-isolation
๋ํ Flash-Attention 2์ ํธํ๋๋ ํ๋์จ์ด๊ฐ ์๋์ง ํ์ธํ์ธ์. flash-attn ์ ์ฅ์์ ๊ณต์ ๋ฌธ์์์ ์์ธํ ์์๋ณด์ธ์. ๋ํ ๋ชจ๋ธ์ ๋ฐ์ ๋ฐ๋(์: torch.float16)๋ก ๋ก๋ํด์ผ ํฉ๋๋ค.
Flash Attention 2๋ฅผ ์ฌ์ฉํ์ฌ ๋ชจ๋ธ์ ๋ก๋ํ๊ณ ์คํํ๋ ค๋ฉด ์๋ ์ฝ๋๋ฅผ ์ฐธ์กฐํ์ธ์:
>>> import torch
>>> import requests
>>> from PIL import Image
>>> from transformers import SiglipProcessor, SiglipModel
>>> device = "cuda" # ๋ชจ๋ธ์ ๋ก๋ํ ์ฅ์น
>>> model = SiglipModel.from_pretrained(
... "google/siglip-so400m-patch14-384",
... attn_implementation="flash_attention_2",
... dtype=torch.float16,
... device_map=device,
... )
>>> processor = SiglipProcessor.from_pretrained("google/siglip-so400m-patch14-384")
>>> url = "http://images.cocodataset.org/val2017/000000039769.jpg"
>>> image = Image.open(requests.get(url, stream=True).raw)
>>> candidate_labels = ["2 cats", "2 dogs"]
# ํ์ดํ๋ผ์ธ ํ๋กฌํํธ ํ
ํ๋ฆฟ์ ๋ฐ๋ผ ๋์ผํ ๊ฒฐ๊ณผ๋ฅผ ์ป์ต๋๋ค
>>> texts = [f'This is a photo of {label}.' for label in candidate_labels]
# ์ค์: ๋ชจ๋ธ์ด ์ด๋ ๊ฒ ํ์ต๋์์ผ๋ฏ๋ก `padding=max_length`๋ฅผ ์ ๋ฌํฉ๋๋ค
>>> inputs = processor(text=texts, images=image, padding="max_length", return_tensors="pt").to(device)
>>> with torch.no_grad():
... with torch.autocast(device):
... outputs = model(**inputs)
>>> logits_per_image = outputs.logits_per_image
>>> probs = torch.sigmoid(logits_per_image) # ์๊ทธ๋ชจ์ด๋ ํ์ฑํ ํจ์๋ฅผ ์ ์ฉํ ํ๋ฅ ์
๋๋ค
>>> print(f"{probs[0][0]:.1%} that image 0 is '{candidate_labels[0]}'")
19.8% that image 0 is '2 cats'
Scaled Dot Product Attention(SDPA) ์ฌ์ฉํ๊ธฐ[using-scaled-dot-product-attention(SDPA)]]
PyTorch๋ torch.nn.functional์ ์ผ๋ถ๋ก ์ค์ผ์ผ๋ ์ ๊ณฑ ์ดํ
์
(SDPA) ์ฐ์ฐ์๋ฅผ ํฌํจํฉ๋๋ค. ์ด ํจ์๋
์
๋ ฅ๊ณผ ์ฌ์ฉ ์ค์ธ ํ๋์จ์ด์ ๋ฐ๋ผ ์ ์ฉํ ์ ์๋ ์ฌ๋ฌ ๊ตฌํ์ ํฌํจํฉ๋๋ค. ์์ธํ ๋ด์ฉ์
๊ณต์ ๋ฌธ์
๋๋ GPU ์ถ๋ก
ํ์ด์ง๋ฅผ ์ฐธ์กฐํ์ธ์.
from_pretrained()์์ attn_implementation="sdpa"๋ฅผ ์ค์ ํ์ฌ SDPA๋ฅผ ๋ช
์์ ์ผ๋ก ์์ฒญํ ์ ์์ต๋๋ค. torch>=2.1.1์ด ์ค์น๋์ด ์๋์ง ํ์ธํ์ธ์.
>>> from transformers import SiglipModel
>>> model = SiglipModel.from_pretrained(
... "google/siglip-so400m-patch14-384",
... attn_implementation="sdpa",
... dtype=torch.float16,
... device_map=device,
... )
์ต์์ ์๋ ํฅ์์ ์ํด ๋ชจ๋ธ์ ๋ฐ์ ๋ฐ๋(์: torch.float16 ๋๋ torch.bfloat16)๋ก ๋ก๋ํ๋ ๊ฒ์ด ์ข์ต๋๋ค.
์์ ์๋ ํฅ์expected-speedups
์๋๋ google/siglip-so400m-patch14-384 ์ฒดํฌํฌ์ธํธ๋ฅผ float16 ์ ๋ฐ๋๋ก ์ฌ์ฉํ๋ transformers์ ๋ค์ดํฐ๋ธ ๊ตฌํ๊ณผ Flash Attention 2 / SDPA ๋ฒ์ ์ ๋ชจ๋ธ์ ๋ค์ํ ๋ฐฐ์น ํฌ๊ธฐ๋ก ๋น๊ตํ ์ถ๋ก ์๊ฐ์ ์์ ์๋ ํฅ์ ๋ค์ด์ด๊ทธ๋จ์
๋๋ค.
SiglipConfig
autodoc SiglipConfig
SiglipTextConfig
autodoc SiglipTextConfig
SiglipVisionConfig
autodoc SiglipVisionConfig
SiglipTokenizer
autodoc SiglipTokenizer - build_inputs_with_special_tokens - get_special_tokens_mask - create_token_type_ids_from_sequences - save_vocabulary
SiglipImageProcessor
autodoc SiglipImageProcessor - preprocess
SiglipImageProcessorFast
autodoc SiglipImageProcessorFast - preprocess
SiglipProcessor
autodoc SiglipProcessor
SiglipModel
autodoc SiglipModel - forward - get_text_features - get_image_features
SiglipTextModel
autodoc SiglipTextModel - forward
SiglipVisionModel
autodoc SiglipVisionModel - forward
SiglipForImageClassification
autodoc SiglipForImageClassification - forward

