# -*- coding: utf-8 -*- """ ChartPipeline: 完整可视化生成管道 此脚本执行完整的数据可视化生成流程,依次调用各模块实现从数据到最终图表的转换 """ import os import json import argparse import logging import time from datetime import datetime from importlib import import_module from pathlib import Path from config import ( base_url, api_key, embed_model_path, data_resource_path, topk, text_data_path, text_index_path, color_data_path, color_index_path, image_data_path, image_index_path, image_list_path, image_resource_path ) import random from concurrent.futures import ProcessPoolExecutor # 配置日志 logging.basicConfig( level=logging.INFO, format='%(asctime)s - %(name)s - %(levelname)s - %(message)s' ) logger = logging.getLogger("ChartPipeline") # 模块配置 MODULES = [ { "name": "create_index", "description": "索引创建模块", "input_type": "none", "output_type": "none" }, { "name": "preprocess", "description": "数据预处理模块", "input_type": "json", "output_type": "json" }, { "name": "chart_type_recommender", "description": "图表类型推荐模块", "input_type": "json", "output_type": "json" }, { "name": "datafact_generator", "description": "数据洞察模块", "input_type": "json", "output_type": "json" }, { "name": "title_generator", "description": "标题生成模块", "input_type": "json", "output_type": "json" }, { "name": "color_recommender", "description": "色彩推荐模块", "input_type": "json", "output_type": "json" }, { "name": "image_recommender", "description": "图像推荐模块", "input_type": "json", "output_type": "json" }, { "name": "all", "description": "综合模块(datafact+title+color+image)", "input_type": "json", "output_type": "json" }, { "name": "infographics_generator", "description": "信息图表生成模块", "input_type": "json", "output_type": "json" }, { "name": "asset_regenerator", "description": "图像/图标统一生图后处理模块", "input_type": "svg", "output_type": "svg" }, { "name": "full_image_polisher", "description": "整图GPT Image润色后处理模块", "input_type": "png", "output_type": "png" }, { "name": "chart_engine", "description": "图表模板实现引擎", "input_type": "json", "output_type": "svg" }, { "name": "title_styler", "description": "标题元素生成模块", "input_type": "json", "output_type": "svg" } ] def run_pipeline(input_path, output_path=None, temp_dir=None, modules_to_run=None, threads=None, chart_name=None, chart_only=False, output_png=False): """ 执行完整的图表生成管道 Args: input_path (str): 输入数据文件路径(可以是文件或目录) output_path (str, optional): 输出文件路径(可以是文件或目录),如果为None则原地修改 temp_dir (str, optional): 临时文件目录,默认使用./tmp modules_to_run (list, optional): 要运行的模块列表,默认运行所有模块 threads (int, optional): 处理目录时的并发线程数,仅在input_path为目录时生效 chart_name (str, optional): 指定图表名称,仅对infographics_generator模块有效 chart_only (bool, optional): 只输出 chart SVG(跳过 title/image/layout),仅对 infographics_generator 有效 output_png (bool, optional): 同时输出 PNG,仅对 infographics_generator 有效 """ # try: if chart_name is not None: output_path = os.path.join(output_path, chart_name) import shutil shutil.rmtree(output_path, ignore_errors=True) os.makedirs(output_path, exist_ok=True) print("modules_to_run: ", modules_to_run) # 如果是create_index模块,单独处理 if modules_to_run and 'create_index' in modules_to_run: return run_single_file( input_path=input_path, output_path=output_path, temp_dir=temp_dir, modules_to_run=modules_to_run, chart_name=chart_name, chart_only=chart_only, output_png=output_png, ) input_path = Path(input_path) output_path = Path(output_path) if output_path else input_path temp_dir = Path(temp_dir) if temp_dir else Path("./tmp") if input_path.is_dir(): # 如果指定了输出目录且不同于输入目录,创建输出目录 if output_path != input_path: output_path.mkdir(parents=True, exist_ok=True) # 获取输入目录下所有JSON文件 input_files = list(input_path.glob('*.json')) random.shuffle(input_files) # 随机打乱文件顺序 if threads and threads > 1: # 使用线程池并行处理文件 with ProcessPoolExecutor(max_workers=threads) as executor: futures = [] for input_file in input_files: # 如果是inplace处理,输出路径就是输入路径 if output_path == input_path: output_file = input_file else: # 确保输出文件保持相同的文件名 output_file = output_path / input_file.name future = executor.submit( run_single_file, input_path=input_file, output_path=output_file, temp_dir=temp_dir, modules_to_run=modules_to_run, chart_name=chart_name, chart_only=chart_only, output_png=output_png, ) futures.append(future) # 等待所有任务完成并检查结果 results = [future.result() for future in futures] return all(results) else: print("单线程顺序处理") # 单线程顺序处理 success = True for input_file in input_files: if output_path == input_path: output_file = input_file else: output_file = output_path / input_file.name success &= run_single_file( input_path=input_file, output_path=output_file, temp_dir=temp_dir, modules_to_run=modules_to_run, chart_name=chart_name, chart_only=chart_only, output_png=output_png, ) return success else: # 单文件处理 if output_path == input_path: output_file = input_path else: # 如果指定了不同的输出路径,确保它的父目录存在 output_path.parent.mkdir(parents=True, exist_ok=True) output_file = output_path return run_single_file( input_path=input_path, output_path=output_file, temp_dir=temp_dir, modules_to_run=modules_to_run, chart_name=chart_name, chart_only=chart_only, output_png=output_png, ) # except Exception as e: # logger.error(f"管道执行失败: {str(e)}") # return False def run_single_file(input_path, output_path, temp_dir=None, modules_to_run=None, chart_name=None, chart_only=False, output_png=False, return_timing=False): """ 处理单个文件的管道逻辑 Args: input_path (Path): 输入JSON文件路径 output_path (Path): 输出路径 temp_dir (Path): 临时文件目录,默认为./tmp modules_to_run (list): 要运行的模块列表 chart_name (str, optional): 指定图表名称,仅对infographics_generator模块有效 chart_only (bool, optional): 只输出 chart SVG,仅对 infographics_generator 有效 output_png (bool, optional): 同时输出 PNG,仅对 infographics_generator 有效 """ # try: run_start_time = time.time() module_timings = {} def finish(success: bool): total_seconds = time.time() - run_start_time logger.info( "run_single_file timing: total=%.3fs modules=%s", total_seconds, {k: round(v, 3) for k, v in module_timings.items()}, ) if return_timing: return { "success": success, "total_seconds": total_seconds, "module_seconds": module_timings, } return success def fail_module(module_name: str): logger.error("模块执行失败: %s", module_name) return finish(False) # 确定要运行的模块 if not modules_to_run: modules_to_run = [m["name"] for m in MODULES] # 如果要运行create_index,单独处理并直接返回 if 'create_index' in modules_to_run: logger.info("执行create_index模块") # 依次调用三个模块的create_index for module_type in ['title', 'color', 'image']: if module_type == 'title': print("Running modules.title_generator.create_index") module = import_module('modules.title_generator.create_index') if not Path(text_index_path).exists(): module.process( data=text_data_path, index_path=text_index_path, data_path=text_data_path, embed_model_path=embed_model_path ) elif module_type == 'color': print("Running modules.color_recommender.create_index") module = import_module('modules.color_recommender.create_index') if not Path(color_index_path).exists(): module.main( input=color_data_path, output=color_index_path, embed_model_path=embed_model_path ) else: # image print("Running modules.image_recommender.create_index") module = import_module('modules.image_recommender.create_index') if not Path(image_index_path).exists(): module.main( image_list_path=image_list_path, image_resource_path=image_resource_path, index_path=image_index_path, data_path=image_data_path, embed_model_path=embed_model_path ) total_seconds = time.time() - run_start_time if return_timing: return { "success": True, "total_seconds": total_seconds, "module_seconds": module_timings, } return True # 创建临时目录 if temp_dir is None: temp_dir = Path("./tmp") temp_dir.mkdir(parents=True, exist_ok=True) # 确保输出目录存在 output_path.parent.mkdir(parents=True, exist_ok=True) # 当前输入文件 current_input = input_path # 获取最后一个模块的配置 last_module = [m for m in MODULES if m["name"] in modules_to_run][-1] # 根据最后一个模块的输出类型决定最终输出文件的扩展名 if last_module["name"] in ["chart_engine", "title_styler", "asset_regenerator"]: final_output = output_path.with_suffix('.svg') elif last_module["name"] == "full_image_polisher": final_output = output_path.with_suffix('.png') elif last_module["name"] == "infographics_generator": final_output = output_path.parent / f"{output_path.stem}_final.svg" else: final_output = output_path.with_suffix('.json') # 记录执行过程 # 依次执行各模块 for i, module_config in enumerate([m for m in MODULES if m["name"] in modules_to_run]): module_name = module_config["name"] module_desc = module_config["description"] module_start_time = time.time() # logger.info(f"执行模块 {i+1}/{len(modules_to_run)}: {module_name} - {module_desc}") # 特殊处理all模块,依次执行datafact_generator、title_generator、color_recommender和image_recommender if module_name == "all": # 数据洞察模块 preprocess_module = import_module(f"modules.preprocess.preprocess") if not should_skip_module(module_name, output_path): step_start_time = time.time() preprocess_module.process(input=str(current_input), output=str(output_path)) module_timings["all.preprocess"] = time.time() - step_start_time current_input = output_path datafact_module = import_module("modules.datafact_generator.datafact_generator") if not should_skip_module("datafact_generator", output_path): step_start_time = time.time() datafact_module.process(input=str(current_input), output=str(output_path)) module_timings["all.datafact_generator"] = time.time() - step_start_time current_input = output_path # 标题生成模块 title_module = import_module("modules.title_generator.title_generator") if not should_skip_module("title_generator", output_path): step_start_time = time.time() title_module.process( input=str(current_input), output=str(output_path), base_url=base_url, api_key=api_key, embed_model_path=embed_model_path, topk=topk, data_path=text_data_path, index_path=text_index_path ) module_timings["all.title_generator"] = time.time() - step_start_time current_input = output_path # 色彩推荐模块 color_module = import_module("modules.color_recommender.color_recommender") if not should_skip_module("color_recommender", output_path): step_start_time = time.time() color_module.process( input=str(current_input), output=str(output_path), base_url=base_url, api_key=api_key, embed_model_path=embed_model_path, data_path=color_data_path, index_path=color_index_path ) module_timings["all.color_recommender"] = time.time() - step_start_time current_input = output_path # 图像推荐模块 image_module = import_module("modules.image_recommender.image_recommender") if not should_skip_module("image_recommender", output_path): step_start_time = time.time() image_module.process( input=str(current_input), output=str(output_path), base_url=base_url, api_key=api_key, embed_model_path=embed_model_path, data_path=image_data_path, index_path=image_index_path, resource_path=image_resource_path ) module_timings["all.image_recommender"] = time.time() - step_start_time current_input = output_path # 特殊处理title_generator模块,传入配置参数 elif module_name == "title_generator": module = import_module(f"modules.{module_name}.{module_name}") if not should_skip_module(module_name, output_path): step_start_time = time.time() module.process( input=str(current_input), output=str(output_path), base_url=base_url, api_key=api_key, embed_model_path=embed_model_path, topk=topk, data_path=text_data_path, index_path=text_index_path ) module_timings["title_generator"] = time.time() - step_start_time current_input = output_path elif module_name == "color_recommender": module = import_module(f"modules.{module_name}.{module_name}") if not should_skip_module(module_name, output_path): step_start_time = time.time() module.process( input=str(current_input), output=str(output_path), base_url=base_url, api_key=api_key, embed_model_path=embed_model_path, data_path=color_data_path, index_path=color_index_path ) module_timings["color_recommender"] = time.time() - step_start_time current_input = output_path elif module_name == "image_recommender": module = import_module(f"modules.{module_name}.{module_name}") print("image_recommender") if not should_skip_module(module_name, output_path): print("process") step_start_time = time.time() module.process( input=str(current_input), output=str(output_path), base_url=base_url, api_key=api_key, embed_model_path=embed_model_path, data_path=image_data_path, index_path=image_index_path, resource_path=image_resource_path ) module_timings["image_recommender"] = time.time() - step_start_time current_input = output_path elif module_name == "infographics_generator": module = import_module(f"modules.{module_name}.{module_name}") if not should_skip_module(module_name, output_path): step_start_time = time.time() result = module.process( input=str(current_input), output=str(output_path), base_url=base_url, api_key=api_key, chart_name=chart_name, chart_only=chart_only, output_png=output_png, ) module_timings["infographics_generator"] = time.time() - step_start_time if result is False: return fail_module(module_name) current_input = output_path elif module_name == "asset_regenerator": module = import_module(f"modules.{module_name}.{module_name}") if not should_skip_module(module_name, output_path): ok = module.process( input=str(current_input), output=str(output_path), base_url=base_url, api_key=api_key, ) if not ok: return False current_input = output_path elif module_name == "full_image_polisher": module = import_module(f"modules.{module_name}.{module_name}") if not should_skip_module(module_name, output_path): ok = module.process( input=str(current_input), output=str(output_path), base_url=base_url, api_key=api_key, ) if not ok: return False current_input = output_path elif module_name == "chart_engine": # 输入是JSON,输出是SVG module = import_module(f"modules.{module_name}.{module_name}") svg_output = output_path.with_suffix('.svg') if True: # not svg_output.exists(): step_start_time = time.time() result = module.process(input=str(current_input), output=str(svg_output)) module_timings["chart_engine"] = time.time() - step_start_time if result is False: return fail_module(module_name) current_input = svg_output # 更新为SVG文件作为下一个模块的输入 elif module_name == "title_styler": # 输入是JSON,输出是SVG module = import_module(f"modules.{module_name}.{module_name}") title_svg = output_path.parent / f"{output_path.stem}_title.svg" if not title_svg.exists(): step_start_time = time.time() result = module.process(input=str(current_input), output=str(title_svg)) module_timings["title_styler"] = time.time() - step_start_time if result is False: return fail_module(module_name) current_input = title_svg # 更新为标题SVG作为下一个模块的输入 else: # 普通模块:输入JSON,输出JSON module = import_module(f"modules.{module_name}.{module_name}") if not should_skip_module(module_name, output_path): step_start_time = time.time() result = module.process(input=str(current_input), output=str(output_path)) module_timings[module_name] = time.time() - step_start_time if result is False: return fail_module(module_name) current_input = output_path module_timings[f"{module_name}.total"] = time.time() - module_start_time return finish(True) # except Exception as e: # logger.error(f"文件处理失败 {input_path}: {str(e)}") # return False def should_skip_module(module_name: str, output_path: Path) -> bool: """检查是否需要跳过模块执行""" try: if not output_path.exists(): return False # 检查JSON文件中的特定字段 with open(output_path) as f: data = json.load(f) skip_conditions = { "preprocess": lambda d: "metadata" in d and "data" in d and "variables" in d and "processed" in d, "chart_type_recommender": lambda d: "chart_type" in d, "datafact_generator": lambda d: "datafacts" in d, "title_generator": lambda d: False,#"titles" in d, "color_recommender": lambda d: "colors" in d, "image_recommender": lambda d: "images" in d } if module_name in skip_conditions: return skip_conditions[module_name](data) return False except Exception as e: logger.warning(f"检查跳过条件时出错: {str(e)}") return False def parse_args(): parser = argparse.ArgumentParser() parser.add_argument('--input', type=str, help='Input json file path or directory (default: config.data_resource_path = data_resource_dirs[0])', default=data_resource_path) parser.add_argument('--output', type=str, help='Output json file path', default='output') parser.add_argument('--temp-dir', type=str, default='tmp') parser.add_argument('--modules', type=str, nargs='+', help='Modules to run', default=['infographics_generator']) parser.add_argument('--threads', type=int, help='Number of threads for directory processing', default=1) parser.add_argument('--chart-name', type=str, help='Specific chart name to use for infographics_generator') parser.add_argument('--chart-only', action='store_true', help='Only output the chart SVG; skip title/image/layout (infographics_generator only)') parser.add_argument('--output-png', action='store_true', help='Also write a PNG next to the final SVG (uses rsvg-convert)') args = parser.parse_args() # 如果没有指定input,从data_resource_path随机选择 if args.input is None and 'create_index' not in args.modules: json_files = [f for f in os.listdir(data_resource_path) if f.endswith('.json')] if not json_files: raise ValueError(f"在 {data_resource_path} 目录下没有找到json文件") args.input = os.path.join(data_resource_path, random.choice(json_files)) print(f"随机选择输入文件: {args.input}") args.output = "tmp.json" print(f"使用默认输出文件: {args.output}") return args def main(): args = parse_args() modules_to_run = None if args.modules: modules_to_run = [m.strip() for m in args.modules] ok = run_pipeline( input_path=args.input, output_path=args.output, temp_dir=args.temp_dir, modules_to_run=modules_to_run, threads=args.threads, chart_name=args.chart_name, chart_only=args.chart_only, output_png=args.output_png, ) raise SystemExit(0 if ok else 1) if __name__ == "__main__": main()