55import os
66import time
77import uuid
8- from typing import Generator , List , Optional
8+ from datetime import date , datetime
9+ from decimal import Decimal
10+ from typing import Dict , Generator , List , Optional
911
1012from tqdm import tqdm
1113
@@ -25,6 +27,8 @@ def __init__(self, input_args: InputArgs):
2527 self .input_args : InputArgs = input_args
2628 self .llm : Optional [BaseLLM ] = None
2729 self .summary : SummaryModel = SummaryModel ()
30+ self ._full_field_written_count : int = 0
31+ self ._file_written_count : Dict [str , int ] = {}
2832
2933 def load_data (self ) -> Generator [Data , None , None ]:
3034 """
@@ -252,6 +256,28 @@ def summarize(self, summary: SummaryModel) -> SummaryModel:
252256 new_summary .finish_time = time .strftime ("%Y%m%d_%H%M%S" , time .localtime ())
253257 return new_summary
254258
259+ @staticmethod
260+ def _json_default (value ):
261+ if isinstance (value , Decimal ):
262+ return float (value )
263+ if isinstance (value , (datetime , date )):
264+ return value .isoformat ()
265+ return str (value )
266+
267+ @classmethod
268+ def _json_dumps (cls , value : dict ) -> str :
269+ return json .dumps (value , ensure_ascii = False , default = cls ._json_default )
270+
271+ def _resolve_field_list (self , input_args : InputArgs ) -> Optional [List [str ]]:
272+ if input_args .executor .result_save .full_field_sample_count <= 0 :
273+ return input_args .executor .result_save .field_list
274+
275+ if self ._full_field_written_count < input_args .executor .result_save .full_field_sample_count :
276+ self ._full_field_written_count += 1
277+ return None
278+
279+ return input_args .executor .result_save .field_list
280+
255281 def write_single_data (
256282 self , path : str , input_args : InputArgs , result_info : ResultInfo
257283 ):
@@ -261,20 +287,20 @@ def write_single_data(
261287 # 如果启用 merge 模式,将所有数据写入同一个文件
262288 if input_args .executor .result_save .merge :
263289 f_n = os .path .join (path , "all_results.jsonl" )
290+ if not self ._should_write_to_file (input_args , f_n ):
291+ return
292+ str_json = self ._build_output_json (input_args , result_info )
264293 with open (f_n , "a" , encoding = "utf-8" ) as f :
265- # if input_args.executor.result_save.raw:
266- # str_json = json.dumps(result_info.to_raw_dict(), ensure_ascii=False)
267- # else:
268- # str_json = json.dumps(result_info.to_dict(), ensure_ascii=False)
269- str_json = json .dumps (result_info .to_raw_dict (), ensure_ascii = False )
270294 f .write (str_json + "\n " )
295+ self ._file_written_count [f_n ] = self ._file_written_count .get (f_n , 0 ) + 1
271296 return
272297
273298 if not input_args .executor .result_save .good and not result_info .eval_status :
274299 return
275300
276301 # 用集合记录已经写过的(字段名, label名)组合,避免重复写入
277302 written_labels = set ()
303+ str_json : Optional [str ] = None
278304
279305 # 遍历 eval_details 的第一层(字段名组合),第二层是List[EvalDetail]
280306 for field_name , eval_detail_list in result_info .eval_details .items ():
@@ -314,18 +340,39 @@ def write_single_data(
314340 # 没有点分割,直接在字段文件夹下创建文件
315341 f_n = os .path .join (field_dir , parts [0 ] + ".jsonl" )
316342
343+ if not self ._should_write_to_file (input_args , f_n ):
344+ continue
345+ if str_json is None :
346+ str_json = self ._build_output_json (input_args , result_info )
317347 with open (f_n , "a" , encoding = "utf-8" ) as f :
318- if input_args .executor .result_save .raw :
319- str_json = json .dumps (result_info .to_raw_dict (), ensure_ascii = False )
320- else :
321- str_json = json .dumps (result_info .to_dict (), ensure_ascii = False )
322348 f .write (str_json + "\n " )
349+ self ._file_written_count [f_n ] = self ._file_written_count .get (f_n , 0 ) + 1
350+
351+ def _should_write_to_file (self , input_args : InputArgs , file_path : str ) -> bool :
352+ limit = input_args .executor .result_save .limit
353+ if limit is None :
354+ return True
355+ return self ._file_written_count .get (file_path , 0 ) < limit
356+
357+ def _build_output_json (self , input_args : InputArgs , result_info : ResultInfo ) -> str :
358+ field_list = self ._resolve_field_list (input_args )
359+ if input_args .executor .result_save .raw :
360+ output_data = result_info .to_raw_dict (field_list = field_list )
361+ else :
362+ output_data = result_info .to_dict (field_list = field_list )
363+ return self ._json_dumps (output_data )
323364
324365 def write_summary (self , path : str , input_args : InputArgs , summary : SummaryModel ):
325366 if not input_args .executor .result_save .bad :
326367 return
327368 with open (path + "/summary.json" , "w" , encoding = "utf-8" ) as f :
328- json .dump (summary .to_dict (), f , indent = 4 , ensure_ascii = False )
369+ json .dump (
370+ summary .to_dict (),
371+ f ,
372+ indent = 4 ,
373+ ensure_ascii = False ,
374+ default = self ._json_default ,
375+ )
329376
330377 def get_summary (self ):
331378 return self .summary
0 commit comments