Views
No views yet
The Composed Image Retrieval (CIR) task aims to retrieve target images using a composed query consisting of a reference image and a modified text. Advanced methods often utilize contrastive learning as the optimization objective, which benefits from adequate positive and negative examples. However, the triplet for CIR incurs high manual annotation costs, resulting in limited positive examples. Furthermore, existing methods commonly use in-batch negative sampling, which reduces the negative number available for the model. To address the problem of lack of positives, we propose a data generation method by leveraging a multi-modal large language model to construct triplets for CIR. To introduce more negatives during fine-tuning, we design a two-stage fine-tuning framework for CIR, whose second stage introduces plenty of static representations of negatives to optimize the representation space rapidly. The above two improvements can be effectively stacked and designed to be plug-and-play, easily applied to existing CIR models without changing their original architectures. Extensive experiments and ablation analysis demonstrate that our method effectively scales positives and negatives and achieves state-of-the-art results on both FashionIQ and CIRR datasets. In addition, our methods also perform well in zero-shot composed image retrieval, providing a new CIR solution for the low-resources scenario.


pip3 install -r requirements.txt1project_base_path
2└─── tgcir
3 | train.py
4 | ...
5
6└─── clip4cir
7 | train.py
8 | ...
9
10└─── blip4cir
11 | train.py
12 | ...
13
14└─── blip24cir
15 | train.py
16 | ...
17
18└─── zscir
19 | ...
20
21└─── data # ckpts of the first stage
22 └─── tgcir
23 └─── clip4cir
24 └─── blip4cir
25 └─── blip24cir
26
27└─── mm_data # generated caption data
28 └─── fiq
29 └─── cirr
30 └─── zs
31
32└─── checkpoints # ckpts of the second stage
33 └─── fiq_clip
34 └─── cirr_clip
35 └─── fiq_blip
36 └─── cirr_blip
37 └─── fiq_blip2
38 └─── cirr_blip2
39 └─── fiq_tgcir
40 └─── cirr_tgcir
41
42└─── fashionIQ_dataset
43 └─── captions
44 | cap.dress.test.json
45 | cap.dress.train.json
46 | cap.dress.val.json
47 | cap.extend_*.train.json
48 | ...
49
50 └─── images
51 | B00006M009.jpg
52 | ...
53
54 └─── image_splits
55 | split.dress.test.json
56 | split.dress.train.json
57 | split.dress.val.json
58 | ...
59
60 | optimized_images.json
61
62└─── cirr_dataset
63 └─── train
64 └─── 0
65 | train-10108-0-img0.png
66 | ...
67 ...
68
69 └─── dev
70 | dev-0-0-img0.png
71 | ...
72
73 └─── test1
74 | test1-0-0-img0.png
75 | ...
76
77 └─── cirr
78 └─── captions
79 | cap.rc2.test1.json
80 | cap.rc2.train.json
81 | cap.rc2.val.json
82 | cap.rc2.train.extend_*.json
83
84 └─── image_splits
85 | split.rc2.test1.json
86 | split.rc2.train.json
87 | split.rc2.val.json
88
89 | optimized_images.json1#FashionIQ stage 2
2python3 zscir/deduplicate_images.py --dataset fiq --dataset fashionIQ_dataset
3
4#CIRR stage 2
5python3 zscir/deduplicate_images.py --dataset cirr --dataset cirr_dataset1#FashionIQ
2python3 zscir/captioner_llava.py --cir_data fiq --k 5
3
4#CIRR
5python3 zscir/captioner_llava.py --cir_data cirr --k 10
6
7# out-of-domain
8python3 zscir/captioner_llava.py --cir_data cc --cc_id 0
9python3 zscir/captioner_llava.py --cir_data cc --cc_id 32
10python3 zscir/captioner_llava.py --cir_data cc --cc_id 64
11python3 zscir/captioner_llava.py --cir_data cc --cc_id 96
12python3 zscir/captioner_llava.py --cir_data cc --cc_id 128
13python3 zscir/captioner_llava.py --cir_data cc --cc_id 160
14python3 zscir/captioner_llava.py --cir_data cc --cc_id 1921#FashionIQ
2python3 zscir/srm_utils.py --dataset fiq --data_path fashionIQ_dataset
3
4#CIRR
5python3 zscir/srm_utils.py --dataset cirr --data_path cirr_dataset1# tgcir
2python3 zscir/get_cir_data.py --model tgcir --data fiq --refer --i2i_rank 10000 --i2i_rank_max 20000 --p_list 2
3python3 zscir/get_cir_data.py --model tgcir --data cirr --i2i_rank 10000 --i2i_rank_max 15000
4
5# clip4cir
6python3 zscir/get_cir_data.py --model clip --data fiq --refer --i2i_rank 10000 --i2i_rank_max 20000 --p_list 2 --word_num 4
7python3 zscir/get_cir_data.py --model clip --data cirr --i2i_rank 10000 --i2i_rank_max 15000 --word_num 8
8
9# blip4cir
10python3 zscir/get_cir_data.py --model blip --data fiq --refer --K 3000 --p_list 2
11python3 zscir/get_cir_data.py --model blip --data cirr
12
13# blip24cir
14python3 zscir/get_cir_data.py --model blip2 --data fiq --K 6000 --refer --p_list 2
15python3 zscir/get_cir_data.py --model blip2 --data cirr --refer
16
17# zs
18# In-Domain
19python3 zscir/get_cir_data.py --model zs --data fiq --p_list 2 --word_num 5
20python3 zscir/get_cir_data.py --model zs --data cirr
21# Our-of-Domain
22python3 zscir/get_cir_data.py --model zs --data ccfiq --p_list 2 --word_num 10
23python3 zscir/get_cir_data.py --model zs --data cccirr1#FashionIQ
2python3 clip4cir/train.py --dataset fiq --batch-size 256 --num-epochs 3 \
3--output_path checkpoints/fiq_clip \
4--bank_path checkpoints/fiq_clip/fiq_bank.pth \
5--learning-rate 2e-5 --tau 0.02 \
6--model_path data/clip4cir/fiq_stage1.pt --plus
7
8#CIRR
9python3 clip4cir/train.py --dataset cirr --batch-size 256 --num-epochs 3 \
10--output_path checkpoints/cirr_clip \
11--bank_path checkpoints/cirr_clip/cirr_bank.pth \
12--learning-rate 2e-5 --tau 0.02 \
13--model_path data/clip4cir/cirr_stage1.pt --plus1#FashionIQ
2python3 clip4cir/validate.py --dataset fiq --data_path fashionIQ_dataset \
3--model_path checkpoints/fiq_clip/best.pt
4
5#CIRR
6python3 clip4cir/validate.py --dataset cirr --data_path cirr_dataset \
7--model_path checkpoints/cirr_clip/best.pt1# Generate 2 json files at submission/clip4cir/
2# Then submit them to the test website: https://cirr.cecs.anu.edu.au/test_process
3python3 clip4cir/cirr_test_submission.py --model_path checkpoints/cirr_clip/best.pt \
4--submission-name clip4cir --data_path cirr_dataset 1#FashionIQ
2python3 tgcir/train.py --dataset fiq --batch-size 256 --num-epochs 5 \
3--output_path checkpoints/fiq_tg \
4--bank_path checkpoints/fiq_tg/fiq_bank.pth \
5--learning-rate 2e-5 --tau 0.02 \
6--model_path data/tgcir/fiq_stage1.pt --plus
7
8#CIRR
9python3 tgcir/train.py --dataset cirr --batch-size 256 --num-epochs 5 \
10--output_path checkpoints/cirr_tg \
11--bank_path checkpoints/cirr_tg/cirr_bank.pth \
12--learning-rate 2e-5 --tau 0.01 \
13--model_path data/tgcir/cirr_stage1.pt --plus1#FashionIQ
2python3 tgcir/validate.py --dataset fiq --data_path fashionIQ_dataset \
3--model_path checkpoints/fiq_tg/best.pt
4
5#CIRR
6python3 tgcir/validate.py --dataset cirr --data_path cirr_dataset \
7--model_path checkpoints/cirr_tg/best.pt1# Generate 2 json files at submission/tgcir/
2# Then submit them to the test website: https://cirr.cecs.anu.edu.au/test_process
3python3 tgcir/cirr_test_submission.py --model_path checkpoints/cirr_tg/best.pt \
4--submission-name tgcir --data_path cirr_dataset 1#FashionIQ
2python3 blip4cir/train.py --dataset fiq --batch-size 128 --num-epochs 10 \
3--output_path checkpoints/fiq_blip \
4--bank_path checkpoints/fiq_blip/fiq_bank.pth \
5--learning-rate 5e-6 --tau 0.03 \
6--model_path data/blip4cir/fiq_stage1.pt --plus
7
8#CIRR
9python3 blip4cir/train.py --dataset cirr --batch-size 128 --num-epochs 3 \
10--output_path checkpoints/cirr_blip \
11--bank_path checkpoints/cirr_blip/cirr_bank.pth \
12--learning-rate 6e-6 --tau 0.02 \
13--model_path data/blip4cir/cirr_stage1.pt --plus1#FashionIQ
2python3 blip4cir/validate.py --dataset fiq --data_path fashionIQ_dataset \
3--model_path checkpoints/fiq_blip/best.pt
4
5#CIRR
6python3 blip4cir/validate.py --dataset cirr --data_path cirr_dataset \
7--model_path checkpoints/cirr_blip/best.pt1# Generate 2 json files at submission/blip4cir/
2# Then submit them to the test website: https://cirr.cecs.anu.edu.au/test_process
3python3 blip4cir/cirr_test_submission.py --model_path checkpoints/cirr_blip/best.pt \
4--submission-name blip4cir --data_path cirr_dataset 1#FashionIQ
2python3 blip24cir/train.py --dataset fiq --batch-size 32 --num-epochs 3 \
3--output_path checkpoints/fiq_blip2 \
4--bank_path checkpoints/fiq_blip2/fiq_bank.pth \
5--learning-rate 1e-5 --tau 0.05 \
6--model_path data/blip24cir/fiq_stage1.pt --plus
7
8#CIRR
9python3 blip24cir/train.py --dataset cirr --batch-size 32 --num-epochs 3 \
10--output_path checkpoints/cirr_blip2 \
11--bank_path checkpoints/cirr_blip2/cirr_bank.pth \
12--learning-rate 1e-5 --tau 0.05 \
13--model_path data/blip24cir/cirr_stage1.pt --plus1#FashionIQ
2python3 blip24cir/validate.py --dataset fiq --data_path fashionIQ_dataset \
3--model_path checkpoints/fiq_blip2/best.pt
4
5#CIRR
6python3 blip24cir/validate.py --dataset cirr --data_path cirr_dataset \
7--model_path checkpoints/cirr_blip2/best.pt1# Generate 2 json files at submission/blip24cir/
2# Then submit them to the test website: https://cirr.cecs.anu.edu.au/test_process
3python3 blip24cir/cirr_test_submission.py --model_path checkpoints/cirr_blip2/best.pt \
4--submission-name blip24cir --data_path cirr_dataset1# Out-Of-Domain
2
3#FashionIQ
4#base
5python3 zscir/train.py --dataset fiq --batch-size 48 --num-epochs 10 \
6--output_path checkpoints/fiq_zs_cc_base \
7--learning-rate 2e-6 --tau 0.01 \
8--use_cc
9
10#bank
11python3 zscir/train_bank.py --dataset fiq --batch-size 128 --num-epochs 5 \
12--output_path checkpoints/fiq_zs_cc \
13--learning-rate 2e-6 --tau 0.02 \
14--bank_path checkpoints/fiq_zs_cc/fiq_bank.pth \
15--use_cc --model_path checkpoints/fiq_zs_cc_base/best.pt
16
17#CIRR
18#base
19python3 zscir/train.py --dataset cirr --batch-size 48 --num-epochs 10 \
20--output_path checkpoints/cirr_zs_cc_base \
21--learning-rate 2e-6 --tau 0.01 \
22--use_cc
23
24#bank
25python3 zscir/train_bank.py --dataset cirr --batch-size 128 --num-epochs 5 \
26--output_path checkpoints/cirr_zs_cc \
27--learning-rate 2e-6 --tau 0.02 \
28--bank_path checkpoints/cirr_zs_cc/cirr_bank.pth \
29--use_cc --model_path checkpoints/cirr_zs_cc_base/best.pt
30
31# In-Domain
32#FashionIQ
33#base
34python3 zscir/train.py --dataset fiq --batch-size 48 --num-epochs 10 \
35--output_path checkpoints/fiq_zs_base \
36--learning-rate 2e-6 --tau 0.01
37
38#bank
39python3 zscir/train_bank.py --dataset fiq --batch-size 128 --num-epochs 5 \
40--output_path checkpoints/fiq_zs \
41--learning-rate 2e-6 --tau 0.02 \
42--bank_path checkpoints/fiq_zs/fiq_bank.pth \
43--model_path checkpoints/fiq_zs_base/best.pt
44
45#CIRR
46#base
47python3 zscir/train.py --dataset cirr --batch-size 48 --num-epochs 10 \
48--output_path checkpoints/cirr_zs_base \
49--learning-rate 2e-6 --tau 0.01
50
51#bank
52python3 zscir/train_bank.py --dataset cirr --batch-size 128 --num-epochs 5 \
53--output_path checkpoints/cirr_zs \
54--learning-rate 2e-6 --tau 0.02 \
55--bank_path checkpoints/cirr_zs/cirr_bank_2.pth \
56--model_path checkpoints/cirr_zs_base/best.pt1#FashionIQ
2python3 zscir/validate.py --dataset fiq --data_path fashionIQ_dataset \
3--model_path checkpoints/fiq_zs/best.pt
4
5python3 zscir/validate.py --dataset fiq --data_path fashionIQ_dataset \
6--model_path checkpoints/fiq_zs_cc/best.pt
7
8#CIRR
9python3 zscir/validate.py --dataset cirr --data_path cirr_dataset \
10--model_path checkpoints/cirr_zs/best.pt
11
12python3 zscir/validate.py --dataset cirr --data_path cirr_dataset \
13--model_path checkpoints/cirr_zs_cc/best.pt1# Generate 2 json files at submission/zscir/
2# Then submit them to the test website: https://cirr.cecs.anu.edu.au/test_process
3python3 zscir/cirr_test_submission.py --model_path checkpoints/cirr_zs/best.pt \
4--submission-name zscir --data_path cirr_dataset
5
6python3 zscir/cirr_test_submission.py --model_path checkpoints/cirr_zs_cc/best.pt \
7--submission-name zscir_cc --data_path cirr_dataset1@article{feng2024improving,
2 title={Improving Composed Image Retrieval via Contrastive Learning with Scaling Positives and Negatives},
3 author={Feng, Zhangchi and Zhang, Richong and Nie, Zhijie},
4 journal={arXiv preprint arXiv:2404.11317},
5 year={2024}
6}