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.
Files changed (185) hide show
  1. {flexeval-0.12.2 → flexeval-0.13.1}/PKG-INFO +4 -4
  2. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/chat_dataset/base.py +28 -1
  3. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/chat_dataset/openai_messages.py +37 -4
  4. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/evaluate_chat_response.py +12 -0
  5. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/evaluate_from_data.py +1 -0
  6. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/evaluate_generation.py +3 -0
  7. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/language_model/__init__.py +1 -0
  8. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/language_model/base.py +9 -0
  9. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/language_model/hf_lm.py +81 -36
  10. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/language_model/openai_batch_api.py +2 -2
  11. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/language_model/vllm_model.py +38 -10
  12. flexeval-0.13.1/flexeval/core/language_model/vllm_serve_lm.py +250 -0
  13. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/metric/__init__.py +1 -0
  14. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/metric/base.py +7 -0
  15. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/metric/exact_match.py +7 -3
  16. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/metric/llm_geval_score.py +22 -12
  17. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/metric/llm_label.py +22 -10
  18. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/metric/llm_score.py +24 -4
  19. flexeval-0.13.1/flexeval/core/metric/math.py +94 -0
  20. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/reward_model/pairwise_judge_reward_model.py +17 -14
  21. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/string_processor/__init__.py +1 -0
  22. flexeval-0.13.1/flexeval/core/string_processor/mgsm.py +58 -0
  23. flexeval-0.13.1/flexeval/preset_configs/EvalSetup/ja_generation/jamcqa.jsonnet +93 -0
  24. {flexeval-0.12.2 → flexeval-0.13.1}/pyproject.toml +4 -3
  25. {flexeval-0.12.2 → flexeval-0.13.1}/LICENSE +0 -0
  26. {flexeval-0.12.2 → flexeval-0.13.1}/README.md +0 -0
  27. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/__init__.py +0 -0
  28. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/__init__.py +0 -0
  29. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/chat_dataset/__init__.py +0 -0
  30. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/chat_dataset/chatbot_bench.py +0 -0
  31. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/chat_dataset/chatbot_bench_datasets/README.md +0 -0
  32. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/chat_dataset/chatbot_bench_datasets/mt-en-ref-gpt4.jsonl +0 -0
  33. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/chat_dataset/chatbot_bench_datasets/mt-en.jsonl +0 -0
  34. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/chat_dataset/chatbot_bench_datasets/mt-ja-ref-gpt4.jsonl +0 -0
  35. {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
  36. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/chat_dataset/chatbot_bench_datasets/mt-ja.jsonl +0 -0
  37. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/chat_dataset/chatbot_bench_datasets/rakuda-v2-ja.jsonl +0 -0
  38. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/chat_dataset/chatbot_bench_datasets/vicuna-en-ref-gpt4.jsonl +0 -0
  39. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/chat_dataset/chatbot_bench_datasets/vicuna-en.jsonl +0 -0
  40. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/chat_dataset/chatbot_bench_datasets/vicuna-ja-ref-gpt4.jsonl +0 -0
  41. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/chat_dataset/chatbot_bench_datasets/vicuna-ja.jsonl +0 -0
  42. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/chat_dataset/sacrebleu_dataset.py +0 -0
  43. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/chat_dataset/template_based.py +0 -0
  44. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/eval_setups.py +0 -0
  45. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/evaluate_multiple_choice.py +0 -0
  46. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/evaluate_pairwise.py +0 -0
  47. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/evaluate_perplexity.py +0 -0
  48. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/evaluate_reward_model.py +0 -0
  49. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/few_shot_generator/__init__.py +0 -0
  50. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/few_shot_generator/balanced.py +0 -0
  51. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/few_shot_generator/base.py +0 -0
  52. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/few_shot_generator/fixed.py +0 -0
  53. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/few_shot_generator/rand.py +0 -0
  54. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/generation_dataset/__init__.py +0 -0
  55. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/generation_dataset/base.py +0 -0
  56. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/generation_dataset/sacrebleu_dataset.py +0 -0
  57. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/generation_dataset/template_based.py +0 -0
  58. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/language_model/litellm_api.py +0 -0
  59. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/language_model/openai_api.py +0 -0
  60. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/metric/bleu.py +0 -0
  61. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/metric/char_f1.py +0 -0
  62. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/metric/code_eval.py +0 -0
  63. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/metric/common_prefix_length.py +0 -0
  64. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/metric/common_string_length.py +0 -0
  65. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/metric/correlation.py +0 -0
  66. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/metric/output_length_stats.py +0 -0
  67. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/metric/perspective_api.py +0 -0
  68. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/metric/repetition_count.py +0 -0
  69. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/metric/rouge.py +0 -0
  70. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/metric/sari.py +0 -0
  71. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/metric/substring_match.py +0 -0
  72. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/metric/utils.py +0 -0
  73. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/metric/xer.py +0 -0
  74. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/multiple_choice_dataset/__init__.py +0 -0
  75. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/multiple_choice_dataset/base.py +0 -0
  76. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/multiple_choice_dataset/template_based.py +0 -0
  77. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/pairwise_comparison/__init__.py +0 -0
  78. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/pairwise_comparison/judge/__init__.py +0 -0
  79. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/pairwise_comparison/judge/base.py +0 -0
  80. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/pairwise_comparison/judge/llm_judge.py +0 -0
  81. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/pairwise_comparison/match.py +0 -0
  82. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/pairwise_comparison/match_maker/__init__.py +0 -0
  83. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/pairwise_comparison/match_maker/all_combinations.py +0 -0
  84. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/pairwise_comparison/match_maker/base.py +0 -0
  85. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/pairwise_comparison/match_maker/random_combinations.py +0 -0
  86. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/pairwise_comparison/scorer/__init__.py +0 -0
  87. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/pairwise_comparison/scorer/base.py +0 -0
  88. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/pairwise_comparison/scorer/bradley_terry.py +0 -0
  89. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/pairwise_comparison/scorer/win_rate.py +0 -0
  90. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/prompt_template/__init__.py +0 -0
  91. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/prompt_template/base.py +0 -0
  92. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/prompt_template/jinja2.py +0 -0
  93. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/result_recorder/__init__.py +0 -0
  94. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/result_recorder/base.py +0 -0
  95. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/result_recorder/local_recorder.py +0 -0
  96. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/result_recorder/wandb_recorder.py +0 -0
  97. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/reward_bench_dataset/__init__.py +0 -0
  98. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/reward_bench_dataset/base.py +0 -0
  99. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/reward_bench_dataset/template_based.py +0 -0
  100. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/reward_model/__init__.py +0 -0
  101. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/reward_model/base.py +0 -0
  102. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/reward_model/log_prob.py +0 -0
  103. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/reward_model/sequence_classification.py +0 -0
  104. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/string_processor/aio.py +0 -0
  105. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/string_processor/base.py +0 -0
  106. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/string_processor/last_line.py +0 -0
  107. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/string_processor/lower.py +0 -0
  108. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/string_processor/nfkc.py +0 -0
  109. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/string_processor/regex.py +0 -0
  110. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/string_processor/string_strip.py +0 -0
  111. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/string_processor/template.py +0 -0
  112. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/text_dataset/__init__.py +0 -0
  113. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/text_dataset/base.py +0 -0
  114. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/text_dataset/hf.py +0 -0
  115. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/text_dataset/jsonl.py +0 -0
  116. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/tokenizer/__init__.py +0 -0
  117. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/tokenizer/base.py +0 -0
  118. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/tokenizer/mecab.py +0 -0
  119. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/tokenizer/sacrebleu_tokenizer.py +0 -0
  120. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/tokenizer/tiktoken_tokenizer.py +0 -0
  121. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/tokenizer/transformers_tokenizer.py +0 -0
  122. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/tokenizer/whitespace.py +0 -0
  123. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/tool_parser/__init__.py +0 -0
  124. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/tool_parser/base.py +0 -0
  125. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/utils/__init__.py +0 -0
  126. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/utils/chat_util.py +0 -0
  127. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/utils/data_util.py +0 -0
  128. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/core/utils/jinja2_utils.py +0 -0
  129. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/EvalSetup/code_chat/mbpp_chat.jsonnet +0 -0
  130. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/EvalSetup/code_generation/jhumaneval.jsonnet +0 -0
  131. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/EvalSetup/code_generation/jhumaneval_tab_indent.jsonnet +0 -0
  132. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/EvalSetup/code_generation/mbpp.jsonnet +0 -0
  133. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/EvalSetup/code_generation/mbpp_tab_indent.jsonnet +0 -0
  134. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/EvalSetup/code_generation/openai_humaneval.jsonnet +0 -0
  135. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/EvalSetup/code_generation/openai_humaneval_tab_indent.jsonnet +0 -0
  136. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/EvalSetup/en_chat/mt-en.jsonnet +0 -0
  137. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/EvalSetup/en_chat/vicuna-en.jsonnet +0 -0
  138. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/EvalSetup/en_generation/babi.jsonnet +0 -0
  139. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/EvalSetup/en_generation/commonsense_qa.jsonnet +0 -0
  140. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/EvalSetup/en_generation/gsm8k.jsonnet +0 -0
  141. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/EvalSetup/en_generation/squad_v1.jsonnet +0 -0
  142. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/EvalSetup/en_generation/trivia_qa.jsonnet +0 -0
  143. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/EvalSetup/en_generation/twitter_sentiment.jsonnet +0 -0
  144. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/EvalSetup/en_multiple_choice/arc_challenge.jsonnet +0 -0
  145. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/EvalSetup/en_multiple_choice/arc_easy.jsonnet +0 -0
  146. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/EvalSetup/en_multiple_choice/commonsense_qa_mc.jsonnet +0 -0
  147. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/EvalSetup/en_multiple_choice/hellaswag.jsonnet +0 -0
  148. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/EvalSetup/en_multiple_choice/openbookqa.jsonnet +0 -0
  149. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/EvalSetup/en_multiple_choice/piqa.jsonnet +0 -0
  150. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/EvalSetup/en_multiple_choice/xwinograd_en.jsonnet +0 -0
  151. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/EvalSetup/en_perplexity/tiny_shakespeare.jsonnet +0 -0
  152. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/EvalSetup/ja_chat/aio_chat.jsonnet +0 -0
  153. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/EvalSetup/ja_chat/elyza_tasks_100.jsonnet +0 -0
  154. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/EvalSetup/ja_chat/mgsm_ja_chat.jsonnet +0 -0
  155. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/EvalSetup/ja_chat/mt-ja.jsonnet +0 -0
  156. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/EvalSetup/ja_chat/rakuda-v2-ja.jsonnet +0 -0
  157. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/EvalSetup/ja_chat/vicuna-ja.jsonnet +0 -0
  158. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/EvalSetup/ja_generation/aio.jsonnet +0 -0
  159. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/EvalSetup/ja_generation/jcommonsenseqa.jsonnet +0 -0
  160. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/EvalSetup/ja_generation/jnli.jsonnet +0 -0
  161. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/EvalSetup/ja_generation/jsquad.jsonnet +0 -0
  162. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/EvalSetup/ja_generation/mgsm_ja.jsonnet +0 -0
  163. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/EvalSetup/ja_generation/wrime_pos_neg.jsonnet +0 -0
  164. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/EvalSetup/ja_generation/xlsum_ja.jsonnet +0 -0
  165. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/EvalSetup/ja_multiple_choice/jcommonsenseqa_mc.jsonnet +0 -0
  166. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/EvalSetup/ja_multiple_choice/xwinograd_ja.jsonnet +0 -0
  167. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/EvalSetup/translation/wmt20_en_ja.jsonnet +0 -0
  168. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/EvalSetup/translation/wmt20_ja_en.jsonnet +0 -0
  169. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/EvalSetup/translation_chat/wmt20_en_ja_chat.jsonnet +0 -0
  170. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/EvalSetup/translation_chat/wmt20_ja_en_chat.jsonnet +0 -0
  171. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/Metric/assistant_eval_en_single_turn.jsonnet +0 -0
  172. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/Metric/assistant_eval_ja_single_turn.jsonnet +0 -0
  173. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/Metric/elyza_tasks_100_eval.jsonnet +0 -0
  174. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/PairwiseJudge/assistant_judge_en_single_turn.jsonnet +0 -0
  175. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/preset_configs/PairwiseJudge/assistant_judge_ja_single_turn.jsonnet +0 -0
  176. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/scripts/__init__.py +0 -0
  177. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/scripts/common.py +0 -0
  178. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/scripts/flexeval_file.py +0 -0
  179. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/scripts/flexeval_lm.py +0 -0
  180. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/scripts/flexeval_pairwise.py +0 -0
  181. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/scripts/flexeval_presets.py +0 -0
  182. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/scripts/flexeval_reward.py +0 -0
  183. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/utils/__init__.py +0 -0
  184. {flexeval-0.12.2 → flexeval-0.13.1}/flexeval/utils/hf_utils.py +0 -0
  185. {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.12.2
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.9, !=2.7.*, !=3.0.*, !=3.1.*, !=3.2.*, !=3.3.*, !=3.4.*, !=3.5.*, !=3.6.*, !=3.7.*, !=3.8.*, !=3.13.*
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 (>=0.8.4,<0.9.0) ; extra == "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
- 'type': 'function',
42
- 'function': { ... }
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,
@@ -45,6 +45,7 @@ def evaluate_from_data(
45
45
  references_list=references_list,
46
46
  extra_info_list=extra_info_list,
47
47
  )
48
+ metric.cleanup_resources()
48
49
 
49
50
  metrics_summary_dict.update(metric_result.summary)
50
51
 
@@ -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
 
@@ -4,3 +4,4 @@ from .litellm_api import LiteLLMChatAPI
4
4
  from .openai_api import OpenAIChatAPI, OpenAICompletionAPI
5
5
  from .openai_batch_api import OpenAIChatBatchAPI
6
6
  from .vllm_model import VLLM
7
+ from .vllm_serve_lm import VLLMServeLM
@@ -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
- from typing import Any, Literal, TypeVar
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
- model_kwargs = get_default_model_kwargs(model_kwargs)
199
- if not load_peft:
200
- self.model: PreTrainedModel = AutoModelForCausalLM.from_pretrained(
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 close(self) -> None:
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.close()
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
- from typing import Any, Literal
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.llm = LLM(model, **model_kwargs)
136
-
137
- if model_limit_tokens == "default":
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})"