ChartPipeline / modules /title_styler /title_styler.py
Ray1ee01's picture
Upload folder using huggingface_hub
0db40c8 verified
Raw
History Blame Contribute Delete
19.5 kB
"""
Title Processor - 包装接口
为title generation提供统一的处理接口
"""
import json
import re
import argparse
from typing import Dict, Union, Optional, List
from modules.title_styler.infographic_title_generator import InfographicTitleGenerator
from modules.title_styler.templates import get_all_templates
def _scale_svg_to_fit(svg_content: str, original_width: float, original_height: float,
max_width: int, template_name: str) -> str:
"""
Scale SVG content to fit within max_width.
Args:
svg_content: Original SVG string
original_width: Original width
original_height: Original height
max_width: Target max width
template_name: Template name for logging
Returns:
Scaled SVG string
"""
# Calculate scale factor to fit within max_width
scale = max_width / original_width
# Find the SVG tag
svg_pattern = r'<svg[^>]*>'
svg_match = re.search(svg_pattern, svg_content)
svg_tag = svg_match.group(0)
# Calculate new dimensions
new_width = original_width * scale
new_height = original_height * scale
# Update width and height in SVG tag
new_svg_tag = re.sub(r'width="[^"]*"', f'width="{new_width:.0f}"', svg_tag)
new_svg_tag = re.sub(r'height="[^"]*"', f'height="{new_height:.0f}"', new_svg_tag)
# Replace the SVG tag
svg_content = svg_content.replace(svg_tag, new_svg_tag)
# Add scale transform to all content
svg_tag_end = svg_content.find('>', svg_content.find('<svg')) + 1
before = svg_content[:svg_tag_end]
after = svg_content[svg_tag_end:]
# Wrap remaining content in scaled group
closing_svg = '</svg>'
closing_pos = after.rfind(closing_svg)
content = after[:closing_pos]
svg_content = f'{before}<g transform="scale({scale:.4f})">{content}</g>{closing_svg}'
print(f"✅ 生成完成 (缩放 {scale:.2f}x): {template_name} ({new_width:.0f}x{new_height:.0f}px)")
return svg_content
def _build_meta_from_result(result: Dict, scaled: bool = False, scale: float = 1.0) -> Dict:
"""Extract a JSON-serialisable metadata bundle from a generator result.
Caller can persist this in info.json so the rendered title is fully
reproducible without re-running the LLM/styler."""
return {
'template_name': result.get('template_name'),
'template_description': result.get('template_description'),
'alignment': result.get('alignment'),
'width': result.get('width'),
'height': result.get('height'),
'segments': result.get('segments') or [],
'split_method': result.get('split_method'),
'primary_color': result.get('primary_color'),
'secondary_color': result.get('secondary_color'),
'background_color': result.get('background_color'),
'scaled_to_fit': bool(scaled),
'scale_factor': scale,
}
def _select_best_result(results: List[Dict], max_width: int, min_scale: float = 0.6,
return_meta: bool = False):
"""
Select the best result from a list of results based on max_width constraint.
Args:
results: List of result dicts with 'width', 'height', 'svg', 'template_name'
max_width: Maximum width constraint
min_scale: Minimum acceptable scale factor (default 0.6). Results requiring
smaller scale will be discarded.
return_meta: When True, return ``(svg_str, meta_dict)`` instead of bare SVG
string. Backward-compatible default keeps the old contract.
Returns:
- return_meta=False (default): Best SVG content string, or None
- return_meta=True: ``(svg_str_or_None, meta_dict_or_None)``
"""
# Filter results that fit within max_width
valid_results = [r for r in results if r['width'] <= max_width]
if not valid_results:
# Use the result with smallest width and scale it to fit max_width
best_result = min(results, key=lambda x: x['width'])
scale = max_width / best_result['width']
# Discard if scale is too small
if scale < min_scale:
print(f"⚠️ 丢弃结果: {best_result['template_name']} 需要缩放 {scale:.2f}x (< {min_scale})")
return (None, None) if return_meta else None
svg = _scale_svg_to_fit(
best_result['svg'],
best_result['width'],
best_result['height'],
max_width,
best_result['template_name']
)
if return_meta:
return svg, _build_meta_from_result(best_result, scaled=True, scale=scale)
return svg
else:
# Select the result with largest width that fits
best_result = max(valid_results, key=lambda x: x['width'])
print(f"✅ 生成完成: {best_result['template_name']} ({best_result['width']:.0f}x{best_result['height']:.0f}px)")
if return_meta:
return best_result['svg'], _build_meta_from_result(best_result)
return best_result['svg']
def process(
input: str = None,
output: str = None,
input_data: Dict = None,
max_width: int = 500,
text_align: str = "left",
background_color: str = "#FFFFFF",
dark: bool = False,
show_embellishment: bool = True,
show_sub_title: bool = True,
font_family: str = None,
return_meta: bool = False,
):
"""
Process function for generating styled title SVG from input data.
Args:
input (str, optional): Path to the input JSON file.
output (str, optional): Path to the output SVG file (if provided, will save to file).
input_data (Dict, optional): Input data dictionary (alternative to file input).
max_width (int, optional): Maximum width constraint for the title. Defaults to 500.
text_align (str, optional): Text alignment. Options: "left", "center", "right". Defaults to "left".
background_color (str, optional): Background color. Defaults to "#FFFFFF".
dark (bool, optional): Whether to use dark mode. Defaults to False.
show_embellishment (bool, optional): Whether to show embellishments (currently unused). Defaults to True.
show_sub_title (bool, optional): Whether to show the subtitle. Defaults to True.
font_family (str, optional): Font family override (currently unused). Defaults to None.
style (str, optional): Title style: normal, comic, simple, professional, all. Defaults to "normal".
Returns:
str: Always returns the generated SVG content as a string.
If output path is provided, also saves to file.
Input JSON Format:
{
"title": "Main title text",
"subtitle": "Subtitle text (optional)",
"primary_color": "#2E7D32",
"secondary_color": "#4CAF50",
"background_color": "#FFFFFF"
}
"""
try:
# Load the data object
if input_data is None:
if input is None:
print("❌ Error: Either input file path or input_data must be provided")
return None
with open(input, 'r', encoding='utf-8') as f:
data = json.load(f)
else:
data = input_data
# Extract data fields
title = data.get('titles').get('main_title')
if not title:
print("❌ Error: 'title' field is required in input data")
return None
subtitle = data.get('titles').get('sub_title') if show_sub_title else None
if not dark:
primary_color = data.get('colors').get('other').get('primary')
secondary_color = data.get('colors').get('other').get('secondary')
else:
primary_color = data.get('colors_dark').get('other').get('primary')
secondary_color = data.get('colors_dark').get('other').get('secondary')
# Handle style parameter: map "Comics" to "comic", otherwise filter out comic templates
if font_family == "Comics":
filter_mode = "comic_only"
else:
filter_mode = "non_comic"
# Create generator with custom template filtering
all_templates = get_all_templates()
if filter_mode == "comic_only":
# Only use comic templates
filtered_templates = [t for t in all_templates if t.style == 'comic']
else:
# Use all non-comic templates
filtered_templates = [t for t in all_templates if t.style != 'comic']
if not filtered_templates:
print("❌ Error: No templates available after filtering")
return None
# Create generator instance
generator = InfographicTitleGenerator(use_llm=True)
# Override templates with filtered ones
generator.templates = filtered_templates
# Generate title with LLM, top_k=1
results = generator.generate(
title=title,
description=subtitle,
primary_color=primary_color,
secondary_color=secondary_color,
background_color=background_color,
max_width=max_width,
alignment=text_align,
top_k=1,
style="comic" if filter_mode == "comic_only" else None
)
if not results:
print("❌ Error: No results generated")
return (None, None) if return_meta else None
# Select best result based on max_width constraint
if return_meta:
svg_content, meta = _select_best_result(results, max_width, return_meta=True)
else:
svg_content = _select_best_result(results, max_width)
meta = None
# Output handling: save to file if output path is provided
if output and svg_content:
with open(output, 'w', encoding='utf-8') as f:
f.write(svg_content)
print(f" 保存: {output}")
if return_meta:
return svg_content, meta
return svg_content
except FileNotFoundError as e:
print(f"❌ Error: Input file not found: {e}")
return (None, None) if return_meta else None
except json.JSONDecodeError as e:
print(f"❌ Error: Invalid JSON format: {e}")
return (None, None) if return_meta else None
except Exception as e:
print(f"❌ Error in title styling: {str(e)}")
import traceback
traceback.print_exc()
return (None, None) if return_meta else None
def process_batch(
input: str = None,
input_data: Dict = None,
max_widths: List[int] = None,
text_align: str = "left",
background_color: str = "#FFFFFF",
dark: bool = False,
show_embellishment: bool = True,
show_sub_title: bool = True,
font_family: str = None,
return_meta: bool = False,
):
"""
Batch process function for generating styled title SVGs with multiple widths.
This function calls LLM only once and generates SVGs for each width.
Args:
input (str, optional): Path to the input JSON file.
input_data (Dict, optional): Input data dictionary (alternative to file input).
max_widths (List[int]): List of maximum width constraints for the titles.
text_align (str, optional): Text alignment. Options: "left", "center", "right". Defaults to "left".
background_color (str, optional): Background color. Defaults to "#FFFFFF".
dark (bool, optional): Whether to use dark mode. Defaults to False.
show_embellishment (bool, optional): Whether to show embellishments (currently unused). Defaults to True.
show_sub_title (bool, optional): Whether to show the subtitle. Defaults to True.
font_family (str, optional): Font family override. Defaults to None.
Returns:
List[str]: List of SVG content strings, one for each max_width in max_widths.
Order matches the order of max_widths.
"""
if max_widths is None or len(max_widths) == 0:
print("❌ Error: max_widths must be provided and non-empty")
return []
try:
# Load the data object
if input_data is None:
if input is None:
print("❌ Error: Either input file path or input_data must be provided")
return []
with open(input, 'r', encoding='utf-8') as f:
data = json.load(f)
else:
data = input_data
# Extract data fields
title = data.get('titles').get('main_title')
if not title:
print("❌ Error: 'title' field is required in input data")
return []
subtitle = data.get('titles').get('sub_title') if show_sub_title else None
if not dark:
primary_color = data.get('colors').get('other').get('primary')
secondary_color = data.get('colors').get('other').get('secondary')
else:
primary_color = data.get('colors_dark').get('other').get('primary')
secondary_color = data.get('colors_dark').get('other').get('secondary')
# Handle style parameter: map "Comics" to "comic", otherwise filter out comic templates
if font_family == "Comics":
filter_mode = "comic_only"
else:
filter_mode = "non_comic"
# Create generator with custom template filtering
all_templates = get_all_templates()
if filter_mode == "comic_only":
filtered_templates = [t for t in all_templates if t.style == 'comic']
else:
filtered_templates = [t for t in all_templates if t.style != 'comic']
if not filtered_templates:
print("❌ Error: No templates available after filtering")
return []
# Create generator instance
generator = InfographicTitleGenerator(use_llm=True)
generator.templates = filtered_templates
# Step 1: Analyze title with LLM (only once)
analysis_result = generator.analyze_title(
title=title,
description=subtitle,
primary_color=primary_color,
background_color=background_color,
alignment=text_align,
top_k=1,
style="comic" if filter_mode == "comic_only" else None
)
if not analysis_result:
print("❌ Error: Title analysis failed")
return []
# Step 2: Generate SVGs for each width (no LLM calls)
results = generator.generate_with_analysis(
analysis_result=analysis_result,
max_widths=max_widths,
secondary_color=secondary_color
)
if not results:
print("❌ Error: No results generated")
return []
# Step 3: Post-process results - match each max_width to best result
# Results are ordered by max_widths, so we can process each
svg_contents = []
metas = []
# Create a mapping from requested max_width to results
results_by_max_width = {}
for result in results:
req_width = result.get('max_width_requested')
if req_width not in results_by_max_width:
results_by_max_width[req_width] = []
results_by_max_width[req_width].append(result)
for max_width in max_widths:
if max_width in results_by_max_width:
width_results = results_by_max_width[max_width]
if return_meta:
svg_content, meta = _select_best_result(width_results, max_width, return_meta=True)
else:
svg_content = _select_best_result(width_results, max_width)
meta = None
else:
# No result for this width — use closest fit + scale.
all_results_list = [r for r in results]
if all_results_list:
if return_meta:
svg_content, meta = _select_best_result(all_results_list, max_width, return_meta=True)
else:
svg_content = _select_best_result(all_results_list, max_width)
meta = None
else:
svg_content = None
meta = None
svg_contents.append(svg_content)
metas.append(meta)
if return_meta:
return list(zip(svg_contents, metas))
return svg_contents
except FileNotFoundError as e:
print(f"❌ Error: Input file not found: {e}")
return []
except json.JSONDecodeError as e:
print(f"❌ Error: Invalid JSON format: {e}")
return []
except Exception as e:
print(f"❌ Error in title styling: {str(e)}")
import traceback
traceback.print_exc()
return []
def main():
"""命令行接口"""
parser = argparse.ArgumentParser(
description='Generate styled title SVG for a chart',
formatter_class=argparse.RawDescriptionHelpFormatter,
epilog="""
Examples:
# Generate with default settings
python title_processor.py -i input.json -o output.svg
# Generate with specific style and width
python title_processor.py -i input.json -o output.svg --style comic --max-width 600
# Generate without subtitle
python title_processor.py -i input.json -o output.svg --no-subtitle
# Generate with center alignment
python title_processor.py -i input.json -o output.svg --text-align center
Input JSON format:
{
"title": "Your Title Here",
"subtitle": "Optional subtitle",
"primary_color": "#2E7D32",
"secondary_color": "#4CAF50",
"background_color": "#FFFFFF"
}
"""
)
parser.add_argument('--input', '-i', type=str, required=True,
help='Input JSON file path')
parser.add_argument('--output', '-o', type=str,
help='Output SVG file path')
parser.add_argument('--max-width', '-w', type=int, default=500,
help='Maximum width constraint for the title (default: 500)')
parser.add_argument('--text-align', '-a', type=str, default='left',
choices=['left', 'center', 'right'],
help='Text alignment: left, center, or right (default: left)')
parser.add_argument('--no-subtitle', action='store_true',
help='Hide the subtitle')
parser.add_argument('--style', '-s', type=str, default='normal',
choices=['normal', 'comic', 'simple', 'professional', 'all'],
help='Title style (default: normal)')
args = parser.parse_args()
svg_content = process(
input=args.input,
output=args.output,
max_width=args.max_width,
text_align=args.text_align,
show_sub_title=not args.no_subtitle,
style=args.style
)
if svg_content:
if args.output:
print("\n✅ Title styling completed successfully.")
else:
# If no output file, print SVG to stdout
print(svg_content)
else:
print("\n❌ Title styling failed.")
exit(1)
if __name__ == '__main__':
main()