(poster_width, poster_height, panels, figures, agent_config)
| 8 | from utils.src.utils import get_json_from_response |
| 9 | |
| 10 | def no_tree_get_layout(poster_width, poster_height, panels, figures, agent_config): |
| 11 | total_input_token, total_output_token = 0, 0 |
| 12 | agent_name = 'ablation_no_tree_layout' |
| 13 | with open(f"prompt_templates/{agent_name}.yaml", "r") as f: |
| 14 | planner_config = yaml.safe_load(f) |
| 15 | |
| 16 | jinja_env = Environment(undefined=StrictUndefined) |
| 17 | template = jinja_env.from_string(planner_config["template"]) |
| 18 | planner_jinja_args = { |
| 19 | 'poster_width': poster_width, |
| 20 | 'poster_height': poster_height, |
| 21 | 'panels': json.dumps(panels, indent=4), |
| 22 | 'figures': json.dumps(figures, indent=4), |
| 23 | } |
| 24 | |
| 25 | planner_model = ModelFactory.create( |
| 26 | model_platform=agent_config['model_platform'], |
| 27 | model_type=agent_config['model_type'], |
| 28 | model_config_dict=agent_config['model_config'], |
| 29 | ) |
| 30 | |
| 31 | planner_agent = ChatAgent( |
| 32 | system_message=planner_config['system_prompt'], |
| 33 | model=planner_model, |
| 34 | message_window_size=None, |
| 35 | ) |
| 36 | |
| 37 | planner_prompt = template.render(**planner_jinja_args) |
| 38 | |
| 39 | num_trials = 0 |
| 40 | |
| 41 | while True: |
| 42 | num_trials += 1 |
| 43 | print(f"Trial {num_trials}: Generating layout...") |
| 44 | planner_agent.reset() |
| 45 | response = planner_agent.step(planner_prompt) |
| 46 | input_token, output_token = account_token(response) |
| 47 | total_input_token += input_token |
| 48 | total_output_token += output_token |
| 49 | |
| 50 | arrangements = get_json_from_response(response.msgs[0].content) |
| 51 | |
| 52 | if len(arrangements) == 0: |
| 53 | print('Error: Empty response, retrying...') |
| 54 | continue |
| 55 | |
| 56 | if not 'panel_arrangement' in arrangements or\ |
| 57 | not 'figure_arrangement' in arrangements or\ |
| 58 | not 'text_arrangement' in arrangements: |
| 59 | print('Error: Invalid response, retrying...') |
| 60 | continue |
| 61 | |
| 62 | if len(arrangements['panel_arrangement']) != len(panels) or\ |
| 63 | len(arrangements['figure_arrangement']) != len(figures): |
| 64 | print('Error: Invalid response, retrying...') |
| 65 | continue |
| 66 | break |
| 67 | return arrangements['panel_arrangement'], arrangements['figure_arrangement'], arrangements['text_arrangement'], input_token, output_token |
no test coverage detected