File size: 2,707 Bytes
58e6885
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
import json
import logging
from pathlib import Path

# 配置日志
logging.basicConfig(
    level=logging.INFO,
    format='%(asctime)s - %(name)s - %(levelname)s - %(message)s'
)
logger = logging.getLogger("CollectTitleData")

def collect_training_data(input_dir: str, output_file: str) -> None:
    """
    收集标题训练数据
    
    Args:
        input_dir (str): 输入数据目录
        output_file (str): 输出JSON文件路径
    """
    try:
        input_path = Path(input_dir)
        output_path = Path(output_file)
        
        # 确保输出目录存在
        output_path.parent.mkdir(parents=True, exist_ok=True)
        
        # 收集所有数据
        training_data = {}
        
        # 处理目录下的所有JSON文件
        for json_file in input_path.glob("*.json"):
            try:
                with open(json_file, "r", encoding="utf-8") as f:
                    data = json.load(f)
                    
                # 使用文件名(不含扩展名)作为chart_id
                chart_id = json_file.stem
                
                # 提取所需字段
                training_data[chart_id] = {
                    "metadata": {
                        "title": data.get("metadata", {}).get("title", ""),
                        "description": data.get("metadata", {}).get("description", ""),
                        "main_insight": data.get("metadata", {}).get("main_insight", "")
                    },
                    "chart_type": data.get("chart_type", []),
                    "datafacts": data.get("datafacts", []),
                    "data": data.get("data", {"columns": [], "data": []})
                }
                
            except Exception as e:
                logger.error(f"处理文件 {json_file} 时出错: {str(e)}")
                continue
        
        # 保存整理后的数据
        with open(output_path, "w", encoding="utf-8") as f:
            json.dump(training_data, f, indent=2, ensure_ascii=False)
            
        logger.info(f"已收集 {len(training_data)} 个图表的数据到 {output_file}")
        
    except Exception as e:
        logger.error(f"收集训练数据失败: {str(e)}")
        raise

if __name__ == "__main__":
    import argparse
    
    parser = argparse.ArgumentParser(description="收集标题生成训练数据")
    parser.add_argument("--input", default="/data/lizhen/input_data/data2",
                      help="输入数据目录")
    parser.add_argument("--output", default="training_data.json",
                      help="输出JSON文件路径")
    
    args = parser.parse_args()
    
    collect_training_data(args.input, args.output)