Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion results/sample_dataset_llm-webkit_evaluation_report.csv
Original file line number Diff line number Diff line change
@@ -1,2 +1,2 @@
extractor,dataset,total_samples,success_rate,overall,code_edit,formula_edit,table_TEDS,table_edit,text_edit
llm-webkit,sample_dataset,4,0.75,0.8667,1.0,1.0,1.0,1.0,0.3333
llm-webkit,sample_dataset,4,0.5,0.9,1.0,1.0,1.0,1.0,0.5
126 changes: 19 additions & 107 deletions results/sample_dataset_llm-webkit_evaluation_results.json

Large diffs are not rendered by default.

8 changes: 4 additions & 4 deletions results/sample_dataset_with_llm-webkit_extraction.jsonl

Large diffs are not rendered by default.

27 changes: 27 additions & 0 deletions webmainbench/data/saver.py
Original file line number Diff line number Diff line change
Expand Up @@ -98,6 +98,9 @@ def save_evaluation_results(results: Union["EvaluationResult", Dict[str, Any]],
else:
results_dict = results

# 移除extracted_content和extracted_content_list字段以减少文件大小
results_dict = DataSaver._remove_content_fields(results_dict)

if format.lower() == "json":
with open(file_path, 'w', encoding='utf-8') as f:
json.dump(results_dict, f, indent=2, ensure_ascii=False)
Expand Down Expand Up @@ -265,6 +268,30 @@ def _save_jsonl_list(data_list: List[Dict[str, Any]], file_path: Union[str, Path
json.dump(item, f, ensure_ascii=False)
f.write('\n')

@staticmethod
def _remove_content_fields(data: Dict[str, Any]) -> Dict[str, Any]:
"""移除extracted_content和extracted_content_list字段以减少保存文件大小"""
import copy

cleaned_data = copy.deepcopy(data)

def remove_fields(obj):
if isinstance(obj, dict):
# 移除extracted_content和extracted_content_list字段
obj.pop('extracted_content', None)
obj.pop('extracted_content_list', None)
# 递归处理嵌套字典和列表
for value in obj.values():
if isinstance(value, (dict, list)):
remove_fields(value)
elif isinstance(obj, list):
for item in obj:
if isinstance(item, (dict, list)):
remove_fields(item)

remove_fields(cleaned_data)
return cleaned_data

@staticmethod
def append_intermediate_results(results: List[Dict[str, Any]],
file_path: Union[str, Path]) -> None:
Expand Down
16 changes: 6 additions & 10 deletions webmainbench/extractors/llm_webkit_extractor.py
Original file line number Diff line number Diff line change
Expand Up @@ -370,20 +370,16 @@ def _load_vllm_model(self):
trust_remote_code=True
)

# vLLM配置
# vLLM配置 - 参考ray_test_qa.py的简化配置
model_kwargs = {
"model": self.inference_config.model_path,
"trust_remote_code": True,
"dtype": self.inference_config.dtype,
"tensor_parallel_size": self.inference_config.tensor_parallel_size,
"max_model_len": self.inference_config.max_tokens,
"max_num_batched_tokens": max(self.inference_config.max_tokens, 8192),
"gpu_memory_utilization": self.inference_config.gpu_memory_utilization,
"enforce_eager": self.inference_config.enforce_eager,
"disable_custom_all_reduce": True,
"load_format": "auto",
}

print(f"🔧 vLLM配置: {model_kwargs}")

self.model = LLM(**model_kwargs)

# 初始化token状态管理器
Expand All @@ -397,8 +393,8 @@ def _load_vllm_model(self):
print("✅ vLLM模型加载成功!")

except Exception as e:
print(f"⚠️ vLLM加载失败,回退到transformers: {e}")
self._load_transformers_model()
print(f"vLLM加载失败: {e}")
raise RuntimeError(f"vLLM模型加载失败: {e}")

def _create_prompt(self, simplified_html: str) -> str:
"""创建分类提示."""
Expand Down Expand Up @@ -463,7 +459,7 @@ def _generate_with_transformers(self, prompt: str) -> str:

except Exception as e:
print(f"⚠️ transformers生成失败: {e}")
return "{}"
raise RuntimeError(f"transformers生成失败: {e}")

def _extract_json_from_text(self, text: str) -> str:
"""从生成的文本中提取JSON"""
Expand Down