Skip to content

Commit a547553

Browse files
authored
feat: image-text (#3)
* fix: loading of in21k vit * fix: config for gh200 * fix: hostfile * fix: hostfile * fix: requirements * fix: multinode * fix: impor tos * fix: model name arg, dataset size, text query * chore: ignore ids to keep * feat: build index and save embeddings + rankings * fix: update for cp medium * fix: update build_map_index_filter * fix: save faster, remove expensive gather * fix: option to save * fix: reqs * fix: protobuf * fix: beam req * fix: beam req * fix: faster multi-gpu index filteirng * feat: filter uuids from json files * fix: print statement * fix: remove duplicates * fix: remove unused * feat: precompute text embeddings * fix: remove unused * fix: updates * fix: download and process * fix: working now * feat: cross attention/coca-esque model * fix: reorder tar so npy near others * feat: add back sampling with replacement (need to test) * feat: train with pretrained * fix: precomputed * feat: cpu offload faster than regular opt w grad checkp * fix: can load and resample now! * feat: sugarcrepe eval * fix: prefix for imagenet * fix: add prefix in trainer for eval * feat: datacomp evals for contrators models * fix: prefix + print after evals * fix: if empty set to none * fix: dataset-size pass via cli * fix: fairness eval * fix: allow null path for imagenet for testing * feat: mlm + contrastive loss * fix: imagenet fixes * fix: deepspeed config * fix: imagenet eval * feat: three towers image <-> text, text <-> frozen * fix: eval steps * fix: hf model updates * fix: vit pos embed * feat: three towers current * fix: multinode fixes * fix: global rank in multinode * fix: progbar only global rank 0 * feat: higher lr * fix: eval strategy epochs logging fix * feat: no clamp logits config * feat: 3 epoch training * feat: update hostfile * fix: 10 epochs * fix: update hostfile * feat: upload embs to atlas * feat: dino v1 * fix: grad check * fix: clip model * feat: 32k vit-l * fix: update hostfile * fix: workers * fix: more logging * fix: no wandb for now * fix: try smaller vit * fix: try more ds stuff * fix: try openclip loss * fix: remove unneeded print * fix: test clip loss * fix: 32k run * fix: are evals broken? * fix: 16k testing * fix: evals * fix: evals * chore: logging * fix: remove prints * feat: config * fix: remove rng, trust openclip * fix: idk? * fix: path * fix: rank * feat: ok now working L14 * feat: 32k higher lr exp * feat: fb vit mae * feat: mae train * fix: map mae * fix: sp * fix: batch size * feat: 10epoch 65k * feat: higher lr * feat: no wd * feat: long train * feat: 81k bs * feat: 3 epoch 65k * feat: 10 epoch * fix: large 3 epoch train * fix: workers * fix: model utils loading * fix: dataloader for datacomp1b * fix: remove pdb * fix: workers * feat: dfn 2b * fix: bs * fix: bs * fix: wandb * fix: imagenet workers * feat: try unidirectional * fix: path for old h100 * fix: map * fix: lets try this again * fix: try fusing * fix: bad code * fix: 32k map fix * fix: bs and default get for dataset * fix: fused * fix; dumb * fix: try this * feat: pos embed with swiglu gated * fix: patch size * fix: runs now * fix: back to mlp * fix: stage 3? * fix: try again * fix: remove pos embed * fix: wtf * feat: mean pool test again? * feat: augments * fix: try no checkpointing * feat: 3 epoch augmentation train * fix: no randaugment * fix: dataset size * feat: 65k run with augs * fix: imagenet path * feat: try resume training multinode * fix: hostfile * fix: no flip for this train * fix: imagenet * refactor: remove unused * refactor: rename text_encoder -> nomic_encoder * refactor: remove captioner * chore: bump pydantic >= 2.0.0 * feat: eval for clip models * feat: v1.5 config * fix: hf code * refactor: move hf tests to separate * chore: remove unused * refactor: remove * refactor: unused code * refactor: not used * fix: remove unused * refactor: remove xattn * refactor: remove xattn * fix: try to resume * fix: v1.5 * fix: remove unused import * fix: remove ema * fix: remove ema * fix: instructions * feat: tracing code * feat: add stacks * feat: export_stacks=True * fix: with_stack * fix: tensorboard profiling (kind of) working * fix: don't profile, test full thing * feat: moar batch * feat: train * refactor: clean up code * feat: download data * fix: pydantic, workers crashing * fix: prefix * chore: ignore data folder * feat: loadable hf model * fix: map pooling bug * fix: comment old pooling * feat: flickr eval running * feat: flickr to config * feat: flickr eval train * fix: flickr eval doesn't hang * feat: biencoder test * fix: enforce no dynamic ntk * feat: unidirectional * feat: base timm models * fix: simplify vit pos_embed * fix: cls token confusion * feat: timm dinov2 with registers * wip vit rotary * feat: yolo 65k scratch vit * fix: hostfile * fix: revert back to bidirectional * fix: spelling * fix: path * fix: wandb * fix: shards * fix: reqs * feat: eva-style models, timm vit-base * fix: timm vit-b 224 image * feat: timm vit-b-16 first experiment * fix: no flip * feat: eva02 vit base * feat: pooling heads from timm vit * feat: add augreg vits as option * fix: remove pooling heads * fix: dumb renaming of model so eva loads with autoconfig * feat: eva config for training * fix: model loading * feat: 65k eva 3 epoch train * feat: map no clamp * fix: hostfile * fix: reduce workers * fix: no clamp * fix: config * feat: v1.5 train * fix: hostfile + config * fix: config for lower lr * fix: hamming * fix: train * feat: hf vision model code * fix: hostfile * fix: path * refactor: clean up code base * refactor: rename * fix: remove hostfile * refactor: remove sugarcrepe * style: black and isort * docs: readme and config fixes * fix: trainers, come back later
1 parent c545be2 commit a547553

84 files changed

Lines changed: 9902 additions & 1844 deletions

Some content is hidden

Large Commits have some content hidden by default. Use the searchbox below for content that may be hidden.

‎.gitignore‎

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,3 +1,5 @@
1+
data/
2+
**/ids_to_keep_*.json
13
*counts.json*
24
medi*.json
35
nq*

‎README.md‎

Lines changed: 15 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -14,10 +14,12 @@
1414
- Huggingface Support for easy loading of common models (Pythia/GPTNeoX, BERT, etc.)
1515
- Masked Language Modeling (MLM) Pretraining
1616
- [Matryoshka Representation Learning](https://arxiv.org/abs/2205.13147) for flexible embedding sizes
17+
- [CLIP](https://arxiv.org/abs/2103.00020) and [LiT](https://arxiv.org/abs/2111.07991) style contrastive learning
18+
- Support for loading popular ViT (e.g. [timm](https://huggingface.co/timm)) models
1719

1820
## Research
1921

20-
* [Nomic Embed: Training a Reproducible Long Context Text Embedder](https://arxiv.org/abs/2402.01613) by Zach Nussbaum, Jack Morris, Andrei Mulyar, and Brandon Duderstadt
22+
* [Nomic Embed: Training a Reproducible Long Context Text Embedder](https://arxiv.org/abs/2402.01613) by Zach Nussbaum, Jack Morris, Andriy Mulyar, and Brandon Duderstadt
2123

2224
## Getting Started and Requirements
2325

@@ -41,7 +43,7 @@ pip3 install torch torchvision torchaudio
4143
Install wheel, packaging, ninja for Flash Attention (so the builds don't take too long)
4244

4345
```bash
44-
pip install wheel packaging ninja
46+
pip install wheel packaging ninja setuptools
4547
```
4648

4749
Install Flash Attention and the custom kernels
@@ -141,6 +143,17 @@ This will train a bert model on all ~200M examples. To change the dataset, you c
141143

142144
To finetune `nomic-bert-embed-v1-unsupervised`, update the config to `configs/train/contrastive_finetune.yaml`.
143145

146+
147+
## Training `nomic-embed-vision-v1.5`
148+
149+
To align a vision model, you will need to curate a large image-text dataset. More details can be found [here](https://github.com/rom1504/img2dataset).
150+
151+
To align `nomic-embed-vision-v1.5` with `nomic-embed-text-v1.5`, you can run the following command:
152+
153+
```bash
154+
deepspeed train.py --deepspeed_config=configs/deepspeed/image_text.json --config=configs/train/nomic_embed_vision_v1.5.yaml --dtype=bf16
155+
```
156+
144157
### Generating Your Own Data
145158

146159
To generate your own data for any step of the pipeline, you can use the provided scripts in `scripts/text`.

‎convert_to_hf.py‎

Lines changed: 20 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -1,24 +1,40 @@
1-
from contrastors.models.huggingface import NomicBertForPreTraining, NomicBertConfig
2-
from contrastors.models.biencoder import BiEncoder, BiEncoderConfig
31
from argparse import ArgumentParser
42

3+
from contrastors.models.biencoder import BiEncoder, BiEncoderConfig
4+
from contrastors.models.dual_encoder import DualEncoder, DualEncoderConfig
5+
from contrastors.models.huggingface import NomicBertConfig, NomicBertForPreTraining, NomicVisionModel
6+
57

68
def parse_args():
79
parser = ArgumentParser()
810
parser.add_argument("--ckpt_path", type=str, required=True)
911
parser.add_argument("--model_name", type=str, required=True)
1012
parser.add_argument("--private", action="store_true")
1113
parser.add_argument("--biencoder", action="store_true")
14+
parser.add_argument("--vision", action="store_true")
1215
return parser.parse_args()
1316

14-
17+
1518
if __name__ == "__main__":
1619
args = parse_args()
1720
if args.biencoder:
1821
config = BiEncoderConfig.from_pretrained(args.ckpt_path)
1922
model = BiEncoder.from_pretrained(args.ckpt_path, config=config)
2023
model = model.trunk
24+
elif args.vision:
25+
NomicBertConfig.register_for_auto_class()
26+
NomicVisionModel.register_for_auto_class("AutoModel")
27+
config = DualEncoderConfig.from_pretrained(args.ckpt_path)
28+
model = DualEncoder.from_pretrained(args.ckpt_path, config=config)
29+
vision = model.vision
30+
hf_config = NomicBertConfig(**model.vision.trunk.config.to_dict())
31+
model = NomicVisionModel(hf_config)
32+
33+
state_dict = vision.state_dict()
34+
state_dict = {k.replace("trunk.", ""): v for k, v in state_dict.items()}
35+
model.load_state_dict(state_dict)
2136
else:
2237
config = NomicBertConfig.from_pretrained(args.ckpt_path)
2338
model = NomicBertForPreTraining.from_pretrained(args.ckpt_path, config=config)
24-
model.push_to_hub(args.model_name, private=args.private)
39+
40+
model.push_to_hub(args.model_name, private=args.private, use_temp_dir=False)

‎requirements.txt‎

Lines changed: 178 additions & 26 deletions
Original file line numberDiff line numberDiff line change
@@ -1,27 +1,179 @@
1-
datasets>=2.16.0
2-
nomic>3.0.0
3-
webdataset
4-
s3fs>=2023.10.0
5-
boto3
6-
google-cloud-storage
7-
wandb
8-
torchmetrics
9-
transformers>=4.34.0
10-
einops
11-
sentencepiece
12-
deepspeed
13-
wheel
14-
packaging
15-
tabulate
16-
av
17-
evaluate
18-
scipy
19-
pydantic<2.0.0
20-
matplotlib
21-
seaborn
22-
tiktoken
23-
openai
24-
mteb
25-
beir
26-
tabulate
1+
accelerate==0.30.1
2+
aiobotocore==2.12.3
3+
aiohttp==3.9.5
4+
aioitertools==0.11.0
5+
aiosignal==1.3.1
6+
annotated-types==0.6.0
7+
anyio==4.3.0
8+
attrs==23.2.0
9+
av==12.0.0
10+
blis==0.7.11
11+
boto3==1.34.69
12+
botocore==1.34.69
13+
braceexpand==0.1.7
14+
cachetools==5.3.3
15+
catalogue==2.0.10
16+
certifi==2024.2.2
17+
charset-normalizer==3.3.2
18+
click==8.1.7
19+
clip-benchmark==1.6.1
20+
cloudpathlib==0.18.1
21+
colorama==0.4.6
22+
confection==0.1.4
23+
contourpy==1.2.1
24+
cycler==0.12.1
25+
cymem==2.0.8
26+
datasets==2.19.1
27+
deepspeed==0.14.2
28+
dill==0.3.8
29+
distro==1.9.0
30+
docker-pycreds==0.4.0
31+
einops==0.8.0
32+
en-core-web-sm @ https://github.com/explosion/spacy-models/releases/download/en_core_web_sm-3.6.0/en_core_web_sm-3.6.0-py3-none-any.whl#sha256=83276fc78a70045627144786b52e1f2728ad5e29e5e43916ec37ea9c26a11212
33+
eval_type_backport==0.2.0
34+
evaluate==0.4.2
35+
filelock==3.14.0
36+
flash-attn==2.5.8
37+
fonttools==4.51.0
38+
frozenlist==1.4.1
39+
fsspec==2024.3.1
40+
ftfy==6.2.0
41+
gitdb==4.0.11
42+
GitPython==3.1.43
43+
google-api-core==2.19.0
44+
google-auth==2.29.0
45+
google-cloud-core==2.4.1
46+
google-cloud-storage==2.16.0
47+
google-crc32c==1.5.0
48+
google-resumable-media==2.7.0
49+
googleapis-common-protos==1.63.0
50+
h11==0.14.0
51+
hjson==3.1.0
52+
httpcore==1.0.5
53+
httpx==0.27.0
54+
huggingface-hub==0.23.0
55+
idna==3.7
56+
iniconfig==2.0.0
57+
Jinja2==3.1.4
58+
jmespath==1.0.1
59+
joblib==1.4.2
60+
jsonlines==4.0.0
61+
kiwisolver==1.4.5
62+
langcodes==3.4.0
63+
language_data==1.2.0
64+
lightning-utilities==0.11.2
65+
littleutils==0.2.2
66+
loguru==0.7.2
67+
marisa-trie==1.1.1
68+
markdown-it-py==3.0.0
69+
MarkupSafe==2.1.5
70+
matplotlib==3.8.4
71+
mdurl==0.1.2
72+
mpmath==1.3.0
73+
mteb==1.8.11
74+
multidict==6.0.5
2775
multiprocess==0.70.15
76+
murmurhash==1.0.10
77+
networkx==3.3
78+
ninja==1.11.1.1
79+
nltk==3.8.1
80+
nomic==3.0.27
81+
numpy==1.24.2
82+
nvidia-cublas-cu12==12.1.3.1
83+
nvidia-cuda-cupti-cu12==12.1.105
84+
nvidia-cuda-nvrtc-cu12==12.1.105
85+
nvidia-cuda-runtime-cu12==12.1.105
86+
nvidia-cudnn-cu12==8.9.2.26
87+
nvidia-cufft-cu12==11.0.2.54
88+
nvidia-curand-cu12==10.3.2.106
89+
nvidia-cusolver-cu12==11.4.5.107
90+
nvidia-cusparse-cu12==12.1.0.106
91+
nvidia-nccl-cu12==2.20.5
92+
nvidia-nvjitlink-cu12==12.4.127
93+
nvidia-nvtx-cu12==12.1.105
94+
ogb==1.3.6
95+
onnx==1.16.0
96+
onnxconverter-common==1.14.0
97+
open-clip-torch==2.24.0
98+
openai==1.28.1
99+
outdated==0.2.2
100+
packaging==24.0
101+
pandas==2.2.2
102+
pathlib_abc==0.1.1
103+
pathy==0.11.0
104+
peft==0.4.0
105+
pillow==10.2.0
106+
platformdirs==4.2.1
107+
pluggy==1.5.0
108+
polars==0.20.25
109+
preshed==3.0.9
110+
pretty-errors==1.2.25
111+
proto-plus==1.23.0
112+
protobuf==3.20.2
113+
psutil==5.9.8
114+
py-cpuinfo==9.0.0
115+
pyarrow==16.0.0
116+
pyarrow-hotfix==0.6
117+
pyasn1==0.6.0
118+
pyasn1_modules==0.4.0
119+
pycocoevalcap==1.2
120+
pycocotools==2.0.7
121+
pydantic==2.7.1
122+
pydantic_core==2.18.2
123+
Pygments==2.18.0
124+
PyJWT==2.8.0
125+
pynvml==11.5.0
126+
pyparsing==3.1.2
127+
pytest==8.2.0
128+
python-dateutil==2.9.0.post0
129+
pytrec-eval-terrier==0.5.6
130+
pytz==2024.1
131+
PyYAML==6.0.1
132+
regex==2024.5.10
133+
requests==2.31.0
134+
rich==13.7.1
135+
rsa==4.9
136+
s3fs==2024.3.1
137+
s3transfer==0.10.1
138+
safetensors==0.4.3
139+
scikit-learn==1.4.2
140+
scipy==1.13.0
141+
seaborn==0.13.2
142+
sentence-transformers==2.7.0
143+
sentencepiece==0.2.0
144+
sentry-sdk==2.1.1
145+
setproctitle==1.3.3
146+
six==1.16.0
147+
smart-open==6.4.0
148+
smmap==5.0.1
149+
sniffio==1.3.1
150+
spacy==3.6.1
151+
spacy-legacy==3.0.12
152+
spacy-loggers==1.0.4
153+
srsly==2.4.8
154+
sympy==1.12
155+
tabulate==0.9.0
156+
thinc==8.1.12
157+
threadpoolctl==3.5.0
158+
tiktoken==0.6.0
159+
timm==1.0.3
160+
tokenizers==0.19.1
161+
torch==2.3.0
162+
torchaudio==2.3.0
163+
torchmetrics==1.4.0
164+
torchvision==0.18.0
165+
tqdm==4.66.4
166+
transformers==4.40.2
167+
triton==2.3.0
168+
typer==0.9.4
169+
typing_extensions==4.11.0
170+
tzdata==2024.1
171+
urllib3==2.2.1
172+
wandb==0.17.0
173+
wasabi==1.1.2
174+
wcwidth==0.2.13
175+
webdataset==0.2.86
176+
wilds==2.0.0
177+
wrapt==1.16.0
178+
xxhash==3.4.1
179+
yarl==1.9.4

‎scripts/image/dataset_size.py‎

Lines changed: 72 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,72 @@
1+
import concurrent.futures
2+
import json
3+
import multiprocessing as mp
4+
from argparse import ArgumentParser
5+
from pathlib import Path
6+
7+
import braceexpand
8+
import fsspec
9+
import pyarrow.parquet as pq
10+
from tqdm import tqdm
11+
12+
13+
def get_dataset_size(shard):
14+
fs = fsspec.filesystem('s3')
15+
try:
16+
with fs.open(shard.replace(".tar", "_stats.json"), "r") as f:
17+
stats = json.load(f)
18+
shard_size = int(stats["successes"])
19+
20+
except Exception as e:
21+
print(f"Error reading {shard}: {e}")
22+
shard_size = 0
23+
24+
return shard_size
25+
26+
27+
if __name__ == "__main__":
28+
parser = ArgumentParser(description="Get the size of a dataset")
29+
parser.add_argument(
30+
"--shards",
31+
type=str,
32+
help="Path to the shards",
33+
default="s3://commonpool-medium/shards/{00000000..00012895}.tar",
34+
)
35+
parser.add_argument("--workers", type=int, help="Number of workers", default=mp.cpu_count())
36+
args = parser.parse_args()
37+
shards = args.shards
38+
39+
shards_list = braceexpand.braceexpand(shards)
40+
shards_list = list(shards_list)
41+
42+
num_shards = len(shards_list)
43+
print(num_shards)
44+
45+
pbar = tqdm(total=num_shards)
46+
47+
total_size = 0
48+
path2size = {}
49+
if args.workers == 1:
50+
for shard in shards_list:
51+
shard_size = get_dataset_size(shard)
52+
path2size[Path(shard).name] = shard_size
53+
total_size += shard_size
54+
pbar.update(1)
55+
else:
56+
with concurrent.futures.ProcessPoolExecutor(max_workers=mp.cpu_count()) as executor:
57+
future2shard = {executor.submit(get_dataset_size, shard): shard for shard in shards_list}
58+
59+
for future in concurrent.futures.as_completed(future2shard):
60+
shard = future2shard[future]
61+
try:
62+
shard_size = future.result()
63+
path2size[Path(shard).name] = shard_size
64+
total_size += shard_size
65+
except Exception as e:
66+
print(f"Shard {shard} generated an exception: {e}")
67+
68+
pbar.update(1)
69+
70+
print(f"Total size: {total_size:,}")
71+
# with open("shard2size.json", "w") as f:
72+
# json.dump(path2size, f, indent=4)

0 commit comments

Comments
 (0)