flexeval 0.12.2__tar.gz → 0.13.1__tar.gz
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- {flexeval-0.12.2 → flexeval-0.13.1}/PKG-INFO +4 -4
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/chat_dataset/base.py +28 -1
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/chat_dataset/openai_messages.py +37 -4
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/evaluate_chat_response.py +12 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/evaluate_from_data.py +1 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/evaluate_generation.py +3 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/language_model/__init__.py +1 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/language_model/base.py +9 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/language_model/hf_lm.py +81 -36
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/language_model/openai_batch_api.py +2 -2
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/language_model/vllm_model.py +38 -10
- flexeval-0.13.1/flexeval/core/language_model/vllm_serve_lm.py +250 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/metric/__init__.py +1 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/metric/base.py +7 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/metric/exact_match.py +7 -3
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/metric/llm_geval_score.py +22 -12
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/metric/llm_label.py +22 -10
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/metric/llm_score.py +24 -4
- flexeval-0.13.1/flexeval/core/metric/math.py +94 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/reward_model/pairwise_judge_reward_model.py +17 -14
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/string_processor/__init__.py +1 -0
- flexeval-0.13.1/flexeval/core/string_processor/mgsm.py +58 -0
- flexeval-0.13.1/flexeval/preset_configs/EvalSetup/ja_generation/jamcqa.jsonnet +93 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/pyproject.toml +4 -3
- {flexeval-0.12.2 → flexeval-0.13.1}/LICENSE +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/README.md +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/__init__.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/__init__.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/chat_dataset/__init__.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/chat_dataset/chatbot_bench.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/chat_dataset/chatbot_bench_datasets/README.md +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/chat_dataset/chatbot_bench_datasets/mt-en-ref-gpt4.jsonl +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/chat_dataset/chatbot_bench_datasets/mt-en.jsonl +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/chat_dataset/chatbot_bench_datasets/mt-ja-ref-gpt4.jsonl +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/chat_dataset/chatbot_bench_datasets/mt-ja-ref-gpt4o-with-human-annotation.jsonl +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/chat_dataset/chatbot_bench_datasets/mt-ja.jsonl +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/chat_dataset/chatbot_bench_datasets/rakuda-v2-ja.jsonl +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/chat_dataset/chatbot_bench_datasets/vicuna-en-ref-gpt4.jsonl +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/chat_dataset/chatbot_bench_datasets/vicuna-en.jsonl +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/chat_dataset/chatbot_bench_datasets/vicuna-ja-ref-gpt4.jsonl +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/chat_dataset/chatbot_bench_datasets/vicuna-ja.jsonl +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/chat_dataset/sacrebleu_dataset.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/chat_dataset/template_based.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/eval_setups.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/evaluate_multiple_choice.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/evaluate_pairwise.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/evaluate_perplexity.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/evaluate_reward_model.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/few_shot_generator/__init__.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/few_shot_generator/balanced.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/few_shot_generator/base.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/few_shot_generator/fixed.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/few_shot_generator/rand.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/generation_dataset/__init__.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/generation_dataset/base.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/generation_dataset/sacrebleu_dataset.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/generation_dataset/template_based.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/language_model/litellm_api.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/language_model/openai_api.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/metric/bleu.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/metric/char_f1.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/metric/code_eval.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/metric/common_prefix_length.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/metric/common_string_length.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/metric/correlation.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/metric/output_length_stats.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/metric/perspective_api.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/metric/repetition_count.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/metric/rouge.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/metric/sari.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/metric/substring_match.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/metric/utils.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/metric/xer.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/multiple_choice_dataset/__init__.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/multiple_choice_dataset/base.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/multiple_choice_dataset/template_based.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/pairwise_comparison/__init__.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/pairwise_comparison/judge/__init__.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/pairwise_comparison/judge/base.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/pairwise_comparison/judge/llm_judge.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/pairwise_comparison/match.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/pairwise_comparison/match_maker/__init__.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/pairwise_comparison/match_maker/all_combinations.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/pairwise_comparison/match_maker/base.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/pairwise_comparison/match_maker/random_combinations.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/pairwise_comparison/scorer/__init__.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/pairwise_comparison/scorer/base.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/pairwise_comparison/scorer/bradley_terry.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/pairwise_comparison/scorer/win_rate.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/prompt_template/__init__.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/prompt_template/base.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/prompt_template/jinja2.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/result_recorder/__init__.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/result_recorder/base.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/result_recorder/local_recorder.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/result_recorder/wandb_recorder.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/reward_bench_dataset/__init__.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/reward_bench_dataset/base.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/reward_bench_dataset/template_based.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/reward_model/__init__.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/reward_model/base.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/reward_model/log_prob.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/reward_model/sequence_classification.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/string_processor/aio.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/string_processor/base.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/string_processor/last_line.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/string_processor/lower.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/string_processor/nfkc.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/string_processor/regex.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/string_processor/string_strip.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/string_processor/template.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/text_dataset/__init__.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/text_dataset/base.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/text_dataset/hf.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/text_dataset/jsonl.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/tokenizer/__init__.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/tokenizer/base.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/tokenizer/mecab.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/tokenizer/sacrebleu_tokenizer.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/tokenizer/tiktoken_tokenizer.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/tokenizer/transformers_tokenizer.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/tokenizer/whitespace.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/tool_parser/__init__.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/tool_parser/base.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/utils/__init__.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/utils/chat_util.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/utils/data_util.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/utils/jinja2_utils.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/EvalSetup/code_chat/mbpp_chat.jsonnet +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/EvalSetup/code_generation/jhumaneval.jsonnet +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/EvalSetup/code_generation/jhumaneval_tab_indent.jsonnet +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/EvalSetup/code_generation/mbpp.jsonnet +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/EvalSetup/code_generation/mbpp_tab_indent.jsonnet +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/EvalSetup/code_generation/openai_humaneval.jsonnet +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/EvalSetup/code_generation/openai_humaneval_tab_indent.jsonnet +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/EvalSetup/en_chat/mt-en.jsonnet +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/EvalSetup/en_chat/vicuna-en.jsonnet +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/EvalSetup/en_generation/babi.jsonnet +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/EvalSetup/en_generation/commonsense_qa.jsonnet +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/EvalSetup/en_generation/gsm8k.jsonnet +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/EvalSetup/en_generation/squad_v1.jsonnet +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/EvalSetup/en_generation/trivia_qa.jsonnet +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/EvalSetup/en_generation/twitter_sentiment.jsonnet +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/EvalSetup/en_multiple_choice/arc_challenge.jsonnet +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/EvalSetup/en_multiple_choice/arc_easy.jsonnet +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/EvalSetup/en_multiple_choice/commonsense_qa_mc.jsonnet +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/EvalSetup/en_multiple_choice/hellaswag.jsonnet +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/EvalSetup/en_multiple_choice/openbookqa.jsonnet +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/EvalSetup/en_multiple_choice/piqa.jsonnet +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/EvalSetup/en_multiple_choice/xwinograd_en.jsonnet +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/EvalSetup/en_perplexity/tiny_shakespeare.jsonnet +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/EvalSetup/ja_chat/aio_chat.jsonnet +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/EvalSetup/ja_chat/elyza_tasks_100.jsonnet +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/EvalSetup/ja_chat/mgsm_ja_chat.jsonnet +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/EvalSetup/ja_chat/mt-ja.jsonnet +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/EvalSetup/ja_chat/rakuda-v2-ja.jsonnet +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/EvalSetup/ja_chat/vicuna-ja.jsonnet +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/EvalSetup/ja_generation/aio.jsonnet +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/EvalSetup/ja_generation/jcommonsenseqa.jsonnet +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/EvalSetup/ja_generation/jnli.jsonnet +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/EvalSetup/ja_generation/jsquad.jsonnet +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/EvalSetup/ja_generation/mgsm_ja.jsonnet +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/EvalSetup/ja_generation/wrime_pos_neg.jsonnet +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/EvalSetup/ja_generation/xlsum_ja.jsonnet +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/EvalSetup/ja_multiple_choice/jcommonsenseqa_mc.jsonnet +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/EvalSetup/ja_multiple_choice/xwinograd_ja.jsonnet +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/EvalSetup/translation/wmt20_en_ja.jsonnet +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/EvalSetup/translation/wmt20_ja_en.jsonnet +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/EvalSetup/translation_chat/wmt20_en_ja_chat.jsonnet +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/EvalSetup/translation_chat/wmt20_ja_en_chat.jsonnet +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/Metric/assistant_eval_en_single_turn.jsonnet +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/Metric/assistant_eval_ja_single_turn.jsonnet +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/Metric/elyza_tasks_100_eval.jsonnet +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/PairwiseJudge/assistant_judge_en_single_turn.jsonnet +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/PairwiseJudge/assistant_judge_ja_single_turn.jsonnet +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/scripts/__init__.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/scripts/common.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/scripts/flexeval_file.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/scripts/flexeval_lm.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/scripts/flexeval_pairwise.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/scripts/flexeval_presets.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/scripts/flexeval_reward.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/utils/__init__.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/utils/hf_utils.py +0 -0
- {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/utils/module_utils.py +0 -0
|
@@ -1,12 +1,11 @@
|
|
|
1
1
|
Metadata-Version: 2.3
|
|
2
2
|
Name: flexeval
|
|
3
|
-
Version: 0.
|
|
3
|
+
Version: 0.13.1
|
|
4
4
|
Summary:
|
|
5
5
|
Author: ryokan-ri
|
|
6
6
|
Author-email: ryokan.ri@sbintuitions.co.jp
|
|
7
|
-
Requires-Python: >=3.
|
|
7
|
+
Requires-Python: >=3.10,<3.13
|
|
8
8
|
Classifier: Programming Language :: Python :: 3
|
|
9
|
-
Classifier: Programming Language :: Python :: 3.9
|
|
10
9
|
Classifier: Programming Language :: Python :: 3.10
|
|
11
10
|
Classifier: Programming Language :: Python :: 3.11
|
|
12
11
|
Classifier: Programming Language :: Python :: 3.12
|
|
@@ -21,6 +20,7 @@ Requires-Dist: jiwer (>=3.0.4,<4.0.0)
|
|
|
21
20
|
Requires-Dist: jsonargparse[jsonnet] (>=4.26.1,<5.0.0)
|
|
22
21
|
Requires-Dist: litellm (>=1.52.9,<2.0.0)
|
|
23
22
|
Requires-Dist: loguru (>=0.7.2,<0.8.0)
|
|
23
|
+
Requires-Dist: math-verify[antlr4-13-2] (>=0.7.0,<0.8.0)
|
|
24
24
|
Requires-Dist: openai (>=1.52.2,<2.0.0)
|
|
25
25
|
Requires-Dist: peft (>=0.10.0,<0.11.0)
|
|
26
26
|
Requires-Dist: pyarrow (==16.1.0)
|
|
@@ -33,7 +33,7 @@ Requires-Dist: smart-open (>=7.1.0,<8.0.0)
|
|
|
33
33
|
Requires-Dist: sudachipy (>=0.6.10)
|
|
34
34
|
Requires-Dist: tiktoken (>=0.9.0,<0.10.0)
|
|
35
35
|
Requires-Dist: transformers[ja,sentencepiece,torch] (>=4.34.1,<5.0.0)
|
|
36
|
-
Requires-Dist: vllm (
|
|
36
|
+
Requires-Dist: vllm (==0.9.2) ; extra == "vllm"
|
|
37
37
|
Requires-Dist: wandb (>=0.17.2,<0.18.0) ; extra == "wandb"
|
|
38
38
|
Description-Content-Type: text/markdown
|
|
39
39
|
|
|
@@ -28,7 +28,34 @@ class ChatInstance:
|
|
|
28
28
|
}
|
|
29
29
|
]
|
|
30
30
|
```
|
|
31
|
-
|
|
31
|
+
|
|
32
|
+
Tool-Calling message must follow the same format as the OpenAI ChatCompletion API.
|
|
33
|
+
https://platform.openai.com/docs/guides/function-calling?api-mode=chat#defining-functions
|
|
34
|
+
```json
|
|
35
|
+
{
|
|
36
|
+
"role": "assistant",
|
|
37
|
+
"content": "content", # `None` is also allowed if `tool_calls` exists.
|
|
38
|
+
"tool_calls": [
|
|
39
|
+
{
|
|
40
|
+
"id": "dummy1",
|
|
41
|
+
"function": {
|
|
42
|
+
"name": "search_web",
|
|
43
|
+
"arguments": "{\"query\": \"flexeval developer\"}" # Note that this is a json string, not a dictionary.
|
|
44
|
+
}
|
|
45
|
+
}
|
|
46
|
+
]
|
|
47
|
+
}
|
|
48
|
+
```
|
|
49
|
+
|
|
50
|
+
The results from tools should be represented as messages with the role "tool":
|
|
51
|
+
```
|
|
52
|
+
{
|
|
53
|
+
"role": "tool",
|
|
54
|
+
"tool_call_id": "dummy1", # Optional, models on OpenAI APIs requires this field.
|
|
55
|
+
"name": "search_web", # Optional, Some HuggingFace models require this field.
|
|
56
|
+
"content": "[{\"title\": \"sbintuitions/flexeval: Flexible evaluation tool...\", \"description\": \"...\"}]",
|
|
57
|
+
}
|
|
58
|
+
""" # noqa: E501
|
|
32
59
|
tools: list[dict[str, Any]] | None = None
|
|
33
60
|
"""
|
|
34
61
|
A list of definitions of tools in the chat.
|
|
@@ -16,6 +16,10 @@ class OpenAIMessagesDataset(ChatDataset):
|
|
|
16
16
|
The difference lies in that this class has 'tool_definition' field, in which
|
|
17
17
|
available tools are listed.
|
|
18
18
|
|
|
19
|
+
Tool-Calling (Function-Calling) is supported in this class.
|
|
20
|
+
It must follow the same format as the OpenAI ChatCompletion API.
|
|
21
|
+
https://platform.openai.com/docs/guides/function-calling?api-mode=chat#defining-functions
|
|
22
|
+
|
|
19
23
|
Parameters:
|
|
20
24
|
file_path (str | list[str] | None): Path or list of paths to `.jsonl` file(s).
|
|
21
25
|
message_key (str): Key used to extract the list of messages from each JSON object.
|
|
@@ -24,25 +28,54 @@ class OpenAIMessagesDataset(ChatDataset):
|
|
|
24
28
|
drop_if_last_from_assistant (bool): If true, when the last utterance is given by assistant, drop it.
|
|
25
29
|
|
|
26
30
|
In Jsonl, each line must have a following structure:
|
|
27
|
-
|
|
31
|
+
```json
|
|
28
32
|
{
|
|
29
33
|
'<message_key>': [
|
|
30
34
|
{
|
|
31
35
|
'role': 'user',
|
|
32
|
-
'content': '
|
|
36
|
+
'content': 'こんにちは。元気が出る言葉を教えて下さい。'
|
|
33
37
|
},
|
|
34
38
|
{
|
|
35
39
|
'role': 'assistant',
|
|
36
40
|
'content': 'こんなのはどうでしょう。どんどんやってください!'
|
|
41
|
+
},
|
|
42
|
+
],
|
|
43
|
+
}
|
|
44
|
+
```
|
|
45
|
+
|
|
46
|
+
Example with tool-calling:
|
|
47
|
+
```json
|
|
48
|
+
{
|
|
49
|
+
'<message_key>': [
|
|
50
|
+
{
|
|
51
|
+
'role': 'user',
|
|
52
|
+
'content': 'こんにちは。元気が出る偉人の言葉を教えて下さい。'
|
|
53
|
+
},
|
|
54
|
+
{
|
|
55
|
+
'role': 'assistant',
|
|
56
|
+
'content': '調べてみますね。',
|
|
57
|
+
'tool_calls': [
|
|
58
|
+
{
|
|
59
|
+
'id': 'dummy1',
|
|
60
|
+
'function': {
|
|
61
|
+
'name': 'web_search',
|
|
62
|
+
'arguments': '{"query": "元気が出る言葉 偉人"}',
|
|
63
|
+
}
|
|
64
|
+
}
|
|
65
|
+
]
|
|
37
66
|
}
|
|
38
67
|
],
|
|
39
68
|
'<tool_definitions_key>': [
|
|
40
69
|
{
|
|
41
|
-
|
|
42
|
-
|
|
70
|
+
"type": "function",
|
|
71
|
+
"function": {
|
|
72
|
+
"name": "web_search",
|
|
73
|
+
...
|
|
74
|
+
}
|
|
43
75
|
}
|
|
44
76
|
]
|
|
45
77
|
}
|
|
78
|
+
```
|
|
46
79
|
"""
|
|
47
80
|
|
|
48
81
|
def __init__(
|
|
@@ -185,6 +185,8 @@ def evaluate_chat_response( # noqa: C901,PLR0912, PLR0915
|
|
|
185
185
|
if mes["content"] is None:
|
|
186
186
|
mes["content"] = ""
|
|
187
187
|
|
|
188
|
+
language_model.cleanup_resources()
|
|
189
|
+
|
|
188
190
|
# Evaluate the generated responses
|
|
189
191
|
metrics_summary_dict: dict[str, float] = {}
|
|
190
192
|
instance_metrics_list: list[dict[str, Any]] = [{} for _ in range(len(all_messages_list))]
|
|
@@ -197,6 +199,7 @@ def evaluate_chat_response( # noqa: C901,PLR0912, PLR0915
|
|
|
197
199
|
for messages, extra_info in zip(all_messages_list, extra_info_list)
|
|
198
200
|
],
|
|
199
201
|
)
|
|
202
|
+
metric.cleanup_resources()
|
|
200
203
|
|
|
201
204
|
metrics_summary_dict.update(metric_result.summary)
|
|
202
205
|
|
|
@@ -217,6 +220,10 @@ def evaluate_chat_response( # noqa: C901,PLR0912, PLR0915
|
|
|
217
220
|
tool_call_validation_result_counter[mes["tool_call_validation_result"]] += 1
|
|
218
221
|
for finish_reason, count in finish_reason_counter.items():
|
|
219
222
|
metrics_summary_dict[f"finish_reason_ratio-{finish_reason}"] = count / sum(finish_reason_counter.values())
|
|
223
|
+
for validation_result, count in tool_call_validation_result_counter.items():
|
|
224
|
+
metrics_summary_dict[f"tool_call_validation_result_ratio-{validation_result}"] = count / sum(
|
|
225
|
+
tool_call_validation_result_counter.values()
|
|
226
|
+
)
|
|
220
227
|
|
|
221
228
|
logger.info(metrics_summary_dict)
|
|
222
229
|
|
|
@@ -229,6 +236,11 @@ def evaluate_chat_response( # noqa: C901,PLR0912, PLR0915
|
|
|
229
236
|
**instance_metrics,
|
|
230
237
|
}
|
|
231
238
|
| ({"raw_lm_output": messages[-1]["raw_content"]} if "raw_content" in messages[-1] else {})
|
|
239
|
+
| (
|
|
240
|
+
{"tool_call_validation_result": messages[-1]["tool_call_validation_result"]}
|
|
241
|
+
if "tool_call_validation_result" in messages[-1]
|
|
242
|
+
else {}
|
|
243
|
+
)
|
|
232
244
|
for messages, references, extra_info, instance_metrics in zip(
|
|
233
245
|
all_messages_list,
|
|
234
246
|
references_list,
|
|
@@ -69,6 +69,8 @@ def evaluate_generation( # noqa: C901
|
|
|
69
69
|
|
|
70
70
|
pbar.update(len(batch))
|
|
71
71
|
|
|
72
|
+
language_model.cleanup_resources()
|
|
73
|
+
|
|
72
74
|
# Evaluate the generated continuations
|
|
73
75
|
metrics_summary_dict: dict[str, float] = {}
|
|
74
76
|
instance_metrics_list: list[dict[str, Any]] = [{} for _ in range(len(eval_instances))]
|
|
@@ -78,6 +80,7 @@ def evaluate_generation( # noqa: C901
|
|
|
78
80
|
references_list=[i.references for i in eval_instances],
|
|
79
81
|
extra_info_list=[i.inputs for i in eval_instances],
|
|
80
82
|
)
|
|
83
|
+
metric.cleanup_resources()
|
|
81
84
|
|
|
82
85
|
metrics_summary_dict.update(metric_result.summary)
|
|
83
86
|
|
|
@@ -240,6 +240,15 @@ class LanguageModel:
|
|
|
240
240
|
return self._batch_compute_chat_log_probs([prompt], [response])[0]
|
|
241
241
|
return self._batch_compute_chat_log_probs(prompt, response)
|
|
242
242
|
|
|
243
|
+
def cleanup_resources(self) -> None:
|
|
244
|
+
"""
|
|
245
|
+
Clean up resources if necessary.
|
|
246
|
+
This method is called when the language model is no longer needed.
|
|
247
|
+
"""
|
|
248
|
+
|
|
249
|
+
def __del__(self) -> None:
|
|
250
|
+
self.cleanup_resources()
|
|
251
|
+
|
|
243
252
|
|
|
244
253
|
def normalize_stop_sequences(
|
|
245
254
|
stop_sequences_list: list[str | list[str] | None],
|
|
@@ -1,7 +1,10 @@
|
|
|
1
1
|
from __future__ import annotations
|
|
2
2
|
|
|
3
3
|
import contextlib
|
|
4
|
-
|
|
4
|
+
import copy
|
|
5
|
+
import gc
|
|
6
|
+
import json
|
|
7
|
+
from typing import Any, Callable, Literal, TypeVar
|
|
5
8
|
|
|
6
9
|
import torch
|
|
7
10
|
import torch.nn.functional as F # noqa: N812
|
|
@@ -138,6 +141,25 @@ def decode_for_lm_continuation(
|
|
|
138
141
|
return entire_text[len(input_text) :]
|
|
139
142
|
|
|
140
143
|
|
|
144
|
+
def deserialize_tool_calls_in_messages(messages: list[dict[str, Any]]) -> None:
|
|
145
|
+
"""
|
|
146
|
+
We adopt the standard OpenAI format, where the 'arguments' field in tool_calls is expected to be a JSON string.
|
|
147
|
+
However, huggingface/transformers expects the 'arguments' field to be a dict.
|
|
148
|
+
https://huggingface.co/docs/transformers/v4.48.2/chat_templating#a-complete-tool-use-example
|
|
149
|
+
|
|
150
|
+
To resolve this mismatch, this function deserializes 'arguments' before passing messages to apply_chat_template.
|
|
151
|
+
|
|
152
|
+
Args:
|
|
153
|
+
messages: A list of messages to deserialize.
|
|
154
|
+
"""
|
|
155
|
+
deserialized_messages = copy.deepcopy(messages)
|
|
156
|
+
for message in deserialized_messages:
|
|
157
|
+
if message["role"] == "assistant" and "tool_calls" in message:
|
|
158
|
+
for item in message["tool_calls"]:
|
|
159
|
+
item["function"]["arguments"] = json.loads(item["function"]["arguments"])
|
|
160
|
+
return deserialized_messages
|
|
161
|
+
|
|
162
|
+
|
|
141
163
|
class HuggingFaceLM(LanguageModel):
|
|
142
164
|
"""
|
|
143
165
|
LanguageModel implementation using Hugging Face Transformers.
|
|
@@ -194,45 +216,56 @@ class HuggingFaceLM(LanguageModel):
|
|
|
194
216
|
self.chat_template_kwargs = chat_template_kwargs or {}
|
|
195
217
|
self.add_special_tokens = add_special_tokens
|
|
196
218
|
self.default_gen_kwargs = default_gen_kwargs or {}
|
|
197
|
-
|
|
198
|
-
|
|
199
|
-
|
|
200
|
-
|
|
201
|
-
model,
|
|
202
|
-
**model_kwargs,
|
|
203
|
-
)
|
|
204
|
-
else:
|
|
205
|
-
from peft import AutoPeftModelForCausalLM
|
|
206
|
-
|
|
207
|
-
self.model = AutoPeftModelForCausalLM.from_pretrained(
|
|
208
|
-
model,
|
|
209
|
-
**model_kwargs,
|
|
210
|
-
)
|
|
211
|
-
|
|
212
|
-
self.model.eval()
|
|
213
|
-
|
|
219
|
+
# `self.model` is initialized lazily to avoid unnecessary memory usage.
|
|
220
|
+
self.model: PreTrainedModel | None = None
|
|
221
|
+
self.model_kwargs = get_default_model_kwargs(model_kwargs)
|
|
222
|
+
self.load_peft = load_peft
|
|
214
223
|
self.amp_dtype = amp_dtype
|
|
215
|
-
if model_limit_tokens == "default":
|
|
216
|
-
hf_config = self.model.config.to_dict()
|
|
217
|
-
if "n_positions" in hf_config:
|
|
218
|
-
model_limit_tokens = hf_config["n_positions"]
|
|
219
|
-
elif "max_position_embeddings" in hf_config:
|
|
220
|
-
model_limit_tokens = hf_config["max_position_embeddings"]
|
|
221
|
-
else:
|
|
222
|
-
msg = (
|
|
223
|
-
"`model_limit_tokens` was set to “default”, but the default max_position_embedeings "
|
|
224
|
-
"could not be found in the config. Set it to `None`."
|
|
225
|
-
)
|
|
226
|
-
logger.warning(msg)
|
|
227
224
|
self.model_limit_tokens = model_limit_tokens
|
|
228
225
|
self.tool_parser = tool_parser
|
|
229
|
-
|
|
230
|
-
transformers.set_seed(random_seed)
|
|
231
|
-
|
|
232
|
-
logger.info(f"model device: {self.model.device}")
|
|
233
|
-
logger.info(f"model dtype: {self.model.dtype}")
|
|
234
226
|
logger.info(f"amp_dtype: {amp_dtype}")
|
|
235
227
|
logger.info(f"random seed: {random_seed}")
|
|
228
|
+
transformers.set_seed(random_seed)
|
|
229
|
+
|
|
230
|
+
@staticmethod
|
|
231
|
+
def load_model(method: Callable) -> Callable:
|
|
232
|
+
"""Decorator to load the model lazily."""
|
|
233
|
+
|
|
234
|
+
def wrapper(self: HuggingFaceLM, *args: tuple, **kwargs: dict) -> Callable:
|
|
235
|
+
if self.model is None:
|
|
236
|
+
if not self.load_peft:
|
|
237
|
+
self.model = AutoModelForCausalLM.from_pretrained(
|
|
238
|
+
self._model_name_or_path,
|
|
239
|
+
**self.model_kwargs,
|
|
240
|
+
)
|
|
241
|
+
else:
|
|
242
|
+
from peft import AutoPeftModelForCausalLM
|
|
243
|
+
|
|
244
|
+
self.model = AutoPeftModelForCausalLM.from_pretrained(
|
|
245
|
+
self._model_name_or_path,
|
|
246
|
+
**self.model_kwargs,
|
|
247
|
+
)
|
|
248
|
+
|
|
249
|
+
self.model.eval()
|
|
250
|
+
|
|
251
|
+
if self.model_limit_tokens == "default":
|
|
252
|
+
hf_config = self.model.config.to_dict()
|
|
253
|
+
if "n_positions" in hf_config:
|
|
254
|
+
self.model_limit_tokens = hf_config["n_positions"]
|
|
255
|
+
elif "max_position_embeddings" in hf_config:
|
|
256
|
+
self.model_limit_tokens = hf_config["max_position_embeddings"]
|
|
257
|
+
else:
|
|
258
|
+
msg = (
|
|
259
|
+
"`model_limit_tokens` was set to “default”, but the default max_position_embedeings "
|
|
260
|
+
"could not be found in the config. Set it to `None`."
|
|
261
|
+
)
|
|
262
|
+
logger.warning(msg)
|
|
263
|
+
|
|
264
|
+
logger.info(f"model device: {self.model.device}")
|
|
265
|
+
logger.info(f"model dtype: {self.model.dtype}")
|
|
266
|
+
return method(self, *args, **kwargs)
|
|
267
|
+
|
|
268
|
+
return wrapper
|
|
236
269
|
|
|
237
270
|
def _get_amp_context(self) -> contextlib.AbstractContextManager:
|
|
238
271
|
if self.amp_dtype is None:
|
|
@@ -276,6 +309,7 @@ class HuggingFaceLM(LanguageModel):
|
|
|
276
309
|
return stop_token_ids
|
|
277
310
|
|
|
278
311
|
@torch.inference_mode()
|
|
312
|
+
@load_model
|
|
279
313
|
def _batch_complete_text(
|
|
280
314
|
self,
|
|
281
315
|
text_list: list[str],
|
|
@@ -350,6 +384,7 @@ class HuggingFaceLM(LanguageModel):
|
|
|
350
384
|
lm_outputs.append(LMOutput(text=decoded_text, finish_reason=finish_reason))
|
|
351
385
|
return lm_outputs
|
|
352
386
|
|
|
387
|
+
@load_model
|
|
353
388
|
def _batch_generate_chat_response(
|
|
354
389
|
self,
|
|
355
390
|
chat_messages_list: list[list[dict[str, Any]]],
|
|
@@ -366,7 +401,7 @@ class HuggingFaceLM(LanguageModel):
|
|
|
366
401
|
chat_messages.insert(0, {"role": "system", "content": self.system_message})
|
|
367
402
|
chat_messages_as_string = [
|
|
368
403
|
self.tokenizer.apply_chat_template(
|
|
369
|
-
chat_messages,
|
|
404
|
+
deserialize_tool_calls_in_messages(chat_messages),
|
|
370
405
|
tools=tools,
|
|
371
406
|
tokenize=False,
|
|
372
407
|
add_generation_prompt=True,
|
|
@@ -390,6 +425,7 @@ class HuggingFaceLM(LanguageModel):
|
|
|
390
425
|
return lm_outputs
|
|
391
426
|
|
|
392
427
|
@torch.inference_mode()
|
|
428
|
+
@load_model
|
|
393
429
|
def _batch_compute_log_probs(
|
|
394
430
|
self,
|
|
395
431
|
text_list: list[str],
|
|
@@ -496,6 +532,7 @@ class HuggingFaceLM(LanguageModel):
|
|
|
496
532
|
total_log_probs = (log_prob_of_next * log_prob_mask).sum(dim=-1)
|
|
497
533
|
return total_log_probs.tolist()
|
|
498
534
|
|
|
535
|
+
@load_model
|
|
499
536
|
def _batch_compute_chat_log_probs(
|
|
500
537
|
self, prompt_list: list[list[dict[str, Any]]], response_list: list[dict[str, Any]]
|
|
501
538
|
) -> list[float]:
|
|
@@ -512,6 +549,14 @@ class HuggingFaceLM(LanguageModel):
|
|
|
512
549
|
response_as_string.append(response_as_string_i)
|
|
513
550
|
return self._batch_compute_log_probs(response_as_string, prefix_list=prompt_as_string)
|
|
514
551
|
|
|
552
|
+
def cleanup_resources(self) -> None:
|
|
553
|
+
del self._model
|
|
554
|
+
self._model = None
|
|
555
|
+
gc.collect()
|
|
556
|
+
if torch.cuda.is_available():
|
|
557
|
+
logger.info("Cleaning up CUDA resources...")
|
|
558
|
+
torch.cuda.empty_cache()
|
|
559
|
+
|
|
515
560
|
def __repr__(self) -> str:
|
|
516
561
|
return f"{self.__class__.__name__}(model={self._model_name_or_path!r})"
|
|
517
562
|
|
|
@@ -279,7 +279,7 @@ class OpenAIChatBatchAPI(LanguageModel):
|
|
|
279
279
|
for res in api_responses
|
|
280
280
|
]
|
|
281
281
|
|
|
282
|
-
def
|
|
282
|
+
def cleanup_resources(self) -> None:
|
|
283
283
|
# in case that the program fails before the file is initialized in __init__
|
|
284
284
|
if not hasattr(self, "temp_jsonl_file"):
|
|
285
285
|
return
|
|
@@ -333,7 +333,7 @@ class OpenAIChatBatchAPI(LanguageModel):
|
|
|
333
333
|
return log_probs
|
|
334
334
|
|
|
335
335
|
def __del__(self) -> None:
|
|
336
|
-
self.
|
|
336
|
+
self.cleanup_resources()
|
|
337
337
|
|
|
338
338
|
def __repr__(self) -> str:
|
|
339
339
|
return f"{self.__class__.__name__}(model={self.model})"
|
|
@@ -1,6 +1,7 @@
|
|
|
1
1
|
from __future__ import annotations
|
|
2
2
|
|
|
3
|
-
|
|
3
|
+
import time
|
|
4
|
+
from typing import TYPE_CHECKING, Any, Callable, Literal
|
|
4
5
|
|
|
5
6
|
import torch
|
|
6
7
|
from loguru import logger
|
|
@@ -10,7 +11,10 @@ from flexeval.core.string_processor import StringProcessor
|
|
|
10
11
|
from flexeval.core.tool_parser.base import ToolParser
|
|
11
12
|
|
|
12
13
|
from .base import LanguageModel, LMOutput, normalize_stop_sequences
|
|
13
|
-
from .hf_lm import decode_for_lm_continuation, get_prefix_and_completion_from_chat
|
|
14
|
+
from .hf_lm import decode_for_lm_continuation, deserialize_tool_calls_in_messages, get_prefix_and_completion_from_chat
|
|
15
|
+
|
|
16
|
+
if TYPE_CHECKING:
|
|
17
|
+
from vllm import LLM
|
|
14
18
|
|
|
15
19
|
|
|
16
20
|
def tokenize_text_for_lm_prefix(
|
|
@@ -122,9 +126,6 @@ class VLLM(LanguageModel):
|
|
|
122
126
|
if "max_new_tokens" in self.default_gen_kwargs:
|
|
123
127
|
self.default_gen_kwargs["max_tokens"] = self.default_gen_kwargs.pop("max_new_tokens")
|
|
124
128
|
|
|
125
|
-
# import from vllm here because it is an extra dependency
|
|
126
|
-
from vllm import LLM
|
|
127
|
-
|
|
128
129
|
model_kwargs = model_kwargs or {}
|
|
129
130
|
# automatically set tensor_parallel_size to the number of GPUs
|
|
130
131
|
if "tensor_parallel_size" not in model_kwargs:
|
|
@@ -132,13 +133,28 @@ class VLLM(LanguageModel):
|
|
|
132
133
|
if "enable_chunked_prefill" not in model_kwargs:
|
|
133
134
|
model_kwargs["enable_chunked_prefill"] = True
|
|
134
135
|
model_kwargs["disable_sliding_window"] = True
|
|
135
|
-
self.
|
|
136
|
-
|
|
137
|
-
|
|
138
|
-
model_limit_tokens = self.llm.llm_engine.get_model_config().max_model_len
|
|
136
|
+
self.model_kwargs = model_kwargs
|
|
137
|
+
# `self.llm` is initialized lazily to avoid unnecessary memory usage.
|
|
138
|
+
self.llm: LLM | None = None
|
|
139
139
|
self.model_limit_tokens = model_limit_tokens
|
|
140
140
|
self.tool_parser = tool_parser
|
|
141
141
|
|
|
142
|
+
@staticmethod
|
|
143
|
+
def load_model(method: Callable) -> Callable:
|
|
144
|
+
"""Decorator to load the model lazily."""
|
|
145
|
+
|
|
146
|
+
def wrapper(self: VLLM, *args: tuple, **kwargs: dict) -> Callable:
|
|
147
|
+
if self.llm is None:
|
|
148
|
+
from vllm import LLM
|
|
149
|
+
|
|
150
|
+
self.llm = LLM(self.model_name, **self.model_kwargs)
|
|
151
|
+
if self.model_limit_tokens == "default":
|
|
152
|
+
self.model_limit_tokens = self.llm.llm_engine.get_model_config().max_model_len
|
|
153
|
+
return method(self, *args, **kwargs)
|
|
154
|
+
|
|
155
|
+
return wrapper
|
|
156
|
+
|
|
157
|
+
@load_model
|
|
142
158
|
def _batch_complete_text(
|
|
143
159
|
self,
|
|
144
160
|
text_list: list[str],
|
|
@@ -215,6 +231,7 @@ class VLLM(LanguageModel):
|
|
|
215
231
|
outputs.append(LMOutput(text=decoded_text, finish_reason=finish_reason))
|
|
216
232
|
return outputs
|
|
217
233
|
|
|
234
|
+
@load_model
|
|
218
235
|
def _batch_generate_chat_response(
|
|
219
236
|
self,
|
|
220
237
|
chat_messages_list: list[list[dict[str, Any]]],
|
|
@@ -228,7 +245,7 @@ class VLLM(LanguageModel):
|
|
|
228
245
|
chat_messages.insert(0, {"role": "system", "content": self.system_message})
|
|
229
246
|
chat_messages_as_string = [
|
|
230
247
|
self.tokenizer.apply_chat_template(
|
|
231
|
-
chat_messages,
|
|
248
|
+
deserialize_tool_calls_in_messages(chat_messages),
|
|
232
249
|
tools=tools,
|
|
233
250
|
tokenize=False,
|
|
234
251
|
add_generation_prompt=True,
|
|
@@ -250,6 +267,7 @@ class VLLM(LanguageModel):
|
|
|
250
267
|
|
|
251
268
|
return lm_outputs
|
|
252
269
|
|
|
270
|
+
@load_model
|
|
253
271
|
def _batch_compute_log_probs(
|
|
254
272
|
self, text_list: list[str], prefix_list: list[str] | None = None, stride: int | None = None
|
|
255
273
|
) -> list[float]:
|
|
@@ -329,6 +347,7 @@ class VLLM(LanguageModel):
|
|
|
329
347
|
|
|
330
348
|
return batch_logprobs
|
|
331
349
|
|
|
350
|
+
@load_model
|
|
332
351
|
def _batch_compute_chat_log_probs(
|
|
333
352
|
self, prompt_list: list[list[dict[str, Any]]], response_list: list[dict[str, Any]]
|
|
334
353
|
) -> list[float]:
|
|
@@ -345,5 +364,14 @@ class VLLM(LanguageModel):
|
|
|
345
364
|
response_as_string.append(response_as_string_i)
|
|
346
365
|
return self._batch_compute_log_probs(response_as_string, prefix_list=prompt_as_string)
|
|
347
366
|
|
|
367
|
+
def cleanup_resources(self) -> None:
|
|
368
|
+
from vllm.distributed import cleanup_dist_env_and_memory
|
|
369
|
+
|
|
370
|
+
del self.llm
|
|
371
|
+
logger.info("cleaning up vLLM resources...")
|
|
372
|
+
cleanup_dist_env_and_memory()
|
|
373
|
+
time.sleep(10) # wait for the vLLM server to release resources
|
|
374
|
+
self.llm = None
|
|
375
|
+
|
|
348
376
|
def __repr__(self) -> str:
|
|
349
377
|
return f"VLLM(model_name={self.model_name})"
|