Skip to content

Commit 2f3e425

Browse files
authored
support lmms-eval (#246)
1 parent f7adbb9 commit 2f3e425

11 files changed

Lines changed: 359 additions & 298 deletions

File tree

‎configs/quantization/methods/Awq/awq_w_only_vlm.yml‎

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -18,11 +18,11 @@ calib:
1818
seed: *seed
1919
eval:
2020
eval_pos: [pretrain, fake_quant]
21-
type: img_txt
22-
name: MME
21+
type: vqa
22+
name: mme
2323
download: False
2424
path: MME dataset path
25-
bs: 16
25+
bs: 1
2626
inference_per_block: False
2727
quant:
2828
method: Awq

‎configs/quantization/methods/GPTQ/gptq_w_only_vlm.yml‎

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -18,11 +18,11 @@ calib:
1818
seed: *seed
1919
eval:
2020
eval_pos: [pretrain, fake_quant]
21-
type: img_txt
22-
name: MME
21+
type: vqa
22+
name: mme
2323
download: False
2424
path: MME dataset path
25-
bs: 16
25+
bs: 1
2626
inference_per_block: False
2727
quant:
2828
method: GPTQ

‎configs/quantization/methods/RTN/rtn_w_a_vlm.yml‎

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -6,11 +6,11 @@ model:
66
torch_dtype: auto
77
eval:
88
eval_pos: [pretrain, fake_quant]
9-
type: img_txt
10-
name: MME
9+
type: vqa
10+
name: mme
1111
download: False
1212
path: MME dataset path
13-
bs: 16
13+
bs: 1
1414
inference_per_block: False
1515
quant:
1616
method: RTN

‎configs/quantization/methods/SmoothQuant/smoothquant_w_a_vlm.yml‎

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -18,11 +18,11 @@ calib:
1818
seed: *seed
1919
eval:
2020
eval_pos: [pretrain, fake_quant]
21-
type: img_txt
22-
name: MME
21+
type: vqa
22+
name: mme
2323
download: False
2424
path: MME dataset path
25-
bs: 16
25+
bs: 1
2626
inference_per_block: False
2727
quant:
2828
method: SmoothQuant

‎llmc/__main__.py‎

Lines changed: 18 additions & 16 deletions
Original file line numberDiff line numberDiff line change
@@ -3,20 +3,22 @@
33
import gc
44
import json
55
import os
6+
import sys
67
import time
78

89
import torch
910
import torch.distributed as dist
1011
import yaml
1112
from easydict import EasyDict
13+
from lmms_eval.utils import make_table
1214
from loguru import logger
1315
from torch.distributed import destroy_process_group, init_process_group
1416

1517
from llmc.compression.quantization import *
1618
from llmc.compression.sparsification import *
1719
from llmc.data import BaseDataset
1820
from llmc.eval import (AccuracyEval, HumanEval, PerplexityEval,
19-
TokenConsistencyEval, VLMEval)
21+
TokenConsistencyEval, VQAEval)
2022
from llmc.models import *
2123
from llmc.utils import (check_config, mkdirs, print_important_package_version,
2224
seed_all, update_autoawq_quant_config,
@@ -46,9 +48,9 @@ def main(config):
4648
if config.eval.type == 'acc':
4749
acc_eval = AccuracyEval(config_for_eval)
4850
eval_list.append(acc_eval)
49-
elif config.eval.type == 'img_txt':
50-
acc_eval = VLMEval(config_for_eval)
51-
eval_list.append(acc_eval)
51+
elif config.eval.type == 'vqa':
52+
vqa_eval = VQAEval(config_for_eval)
53+
eval_list.append(vqa_eval)
5254
elif config.eval.type == 'code' and config.eval.name == 'human_eval':
5355
human_eval = HumanEval(model.get_tokenizer(), config_for_eval)
5456
eval_list.append(human_eval)
@@ -61,10 +63,11 @@ def main(config):
6163
for acc_eval in eval_list:
6264
acc = acc_eval.eval(model)
6365
logger.info(f'{config.eval.name} acc : {acc}')
64-
elif config.eval.type == 'img_txt':
65-
for vlm_eval in eval_list:
66-
results = vlm_eval.eval(model)
67-
logger.info(f'{config.eval.name} results : {results}')
66+
elif config.eval.type == 'vqa':
67+
for vqa_eval in eval_list:
68+
results = vqa_eval.eval(model)
69+
logger.info(f'{config.eval.name} results :')
70+
print(make_table(results))
6871
elif config.eval.type == 'code' and config.eval.name == 'human_eval':
6972
for human_eval in eval_list:
7073
results = human_eval.eval(model, eval_pos='pretrain')
@@ -125,10 +128,6 @@ def main(config):
125128
for acc_eval in eval_list:
126129
acc = acc_eval.eval(model)
127130
logger.info(f'{config.eval.name} acc : {acc}')
128-
elif config.eval.type == 'img_txt':
129-
for vlm_eval in eval_list:
130-
results = vlm_eval.eval(model)
131-
logger.info(f'{config.eval.name} results : {results}')
132131
elif config.eval.type == 'code' and config.eval.name == 'human_eval':
133132
for human_eval in eval_list:
134133
results = human_eval.eval(model, eval_pos='transformed')
@@ -157,10 +156,12 @@ def main(config):
157156
for acc_eval in eval_list:
158157
acc = acc_eval.eval(model)
159158
logger.info(f'{config.eval.name} acc : {acc}')
160-
elif config.eval.type == 'img_txt':
161-
for vlm_eval in eval_list:
162-
results = vlm_eval.eval(model)
163-
logger.info(f'{config.eval.name} results : {results}')
159+
160+
elif config.eval.type == 'vqa':
161+
for vqa_eval in eval_list:
162+
results = vqa_eval.eval(model)
163+
logger.info(f'{config.eval.name} results :')
164+
print(make_table(results))
164165
elif config.eval.type == 'code' and config.eval.name == 'human_eval':
165166
for human_eval in eval_list:
166167
results = human_eval.eval(model, eval_pos='fake_quant')
@@ -251,6 +252,7 @@ def main(config):
251252

252253

253254
if __name__ == '__main__':
255+
logger.add(sys.stdout, level='INFO')
254256
llmc_start_time = time.time()
255257
parser = argparse.ArgumentParser()
256258
parser.add_argument('--config', type=str, required=True)

‎llmc/eval/__init__.py‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -2,4 +2,4 @@
22
from .eval_code import HumanEval
33
from .eval_ppl import PerplexityEval
44
from .eval_token_consist import TokenConsistencyEval
5-
from .eval_vlm import VLMEval
5+
from .eval_vqa import VQAEval

‎llmc/eval/eval_base.py‎

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -22,7 +22,8 @@ def __init__(self, tokenizer, config):
2222
'c4',
2323
'ptb',
2424
'custom',
25-
'human_eval'
25+
'human_eval',
26+
'mme',
2627
], 'Eval only support wikitext2, c4, ptb, custom, human_eval dataset now.'
2728
self.seq_len = self.eval_cfg.get('seq_len', None)
2829
self.bs = self.eval_cfg['bs']

0 commit comments

Comments
 (0)