diff --git a/feedback.py b/feedback.py index f0f0f9e..26c6a89 100644 --- a/feedback.py +++ b/feedback.py @@ -164,6 +164,9 @@ class FeedbackDB: "instruction": f"Generate a pixel-art sprite for: {entry.prompt}", "response": entry.params, "rating": entry.rating, + "feedback": entry.feedback, + "image_path": entry.image_path, + "prompt": entry.prompt, } ) lines.append(line) diff --git a/generate_sample_mcp.py b/generate_sample_mcp.py new file mode 100644 index 0000000..3bd802f --- /dev/null +++ b/generate_sample_mcp.py @@ -0,0 +1,85 @@ +#!/usr/bin/env python3 +"""Generate a sample dataset through the MCP server.""" + +import asyncio +import os +import sys + +from mcp.client.session import ClientSession +from mcp.client.stdio import StdioServerParameters, stdio_client + +BASE_DIR = os.path.dirname(os.path.abspath(__file__)) + +SAMPLE_SPECS = [ + { + "prompt": "a brave knight in shining armor holding a sword", + "output_path": "tetra/knight_01.png", + "seed": 42, + }, + { + "prompt": "an elven archer with a longbow and green cloak", + "output_path": "tetra/archer_01.png", + "seed": 43, + }, + { + "prompt": "a robed mage casting a blue fireball", + "output_path": "tetra/mage_01.png", + "seed": 44, + }, + { + "prompt": "a peasant worker carrying a wooden hammer", + "output_path": "tetra/worker_01.png", + "seed": 45, + }, + { + "prompt": "a crystal golem with geometric facets", + "output_path": "tetra/golem_01.png", + "seed": 46, + }, + { + "prompt": "a chaos demon with horns and jagged armor", + "output_path": "tetra/demon_01.png", + "seed": 47, + }, + { + "prompt": "an ancient treant with glowing green eyes", + "output_path": "tetra/treant_01.png", + "seed": 48, + }, + { + "prompt": "a skeleton warrior with a rusted sword", + "output_path": "tetra/skeleton_01.png", + "seed": 49, + }, +] + + +async def main(): + server_path = os.path.join(BASE_DIR, "server.py") + params = StdioServerParameters( + command=sys.executable, + args=[server_path], + env={**os.environ, "PYTHONUNBUFFERED": "1"}, + ) + + async with stdio_client(params) as (read, write): + async with ClientSession(read, write) as session: + await session.initialize() + + print("Calling batch_generate...") + result = await session.call_tool("batch_generate", {"specs": SAMPLE_SPECS}) + + print("\nResults:") + for item in result.content: + if item.type == "text": + print(item.text) + + stats = await session.call_tool("db_stats", {}) + print("\nDB stats:") + for item in stats.content: + if item.type == "text": + print(item.text) + + +if __name__ == "__main__": + asyncio.run(main()) diff --git a/review_gui.py b/review_gui.py index 8c9c706..3479087 100644 --- a/review_gui.py +++ b/review_gui.py @@ -28,6 +28,8 @@ _texture_registry = None def _load_thumbnail(path, tag): if not path or not os.path.exists(path): return False + if dpg.does_item_exist(tag): + return True try: img = Image.open(path).convert("RGBA") img.thumbnail((THUMB_SIZE, THUMB_SIZE), Image.NEAREST) @@ -108,54 +110,90 @@ def _build_entry_row(entry): dpg.add_text(p, color=(100, 100, 100)) dpg.add_spacer(height=4) - with dpg.group(horizontal=True): - for star in range(1, 6): - filled = star <= entry.rating - label = "\u2605" if filled else "\u2606" - color = (255, 200, 80) if filled else (80, 80, 90) - btn_tag = f"star_{entry.id}_{star}" - dpg.add_button( - label=label, - tag=btn_tag, - callback=lambda s, a, e=entry: _rate(e, s), - width=28, - height=24, - ) - dpg.bind_item_theme(btn_tag, _star_theme(color)) + with dpg.table( + header_row=False, + policy=dpg.mvTable_SizingFixedFit, + no_host_extendX=True, + row_background=False, + borders_innerH=False, + borders_outerH=False, + borders_innerV=False, + borders_outerV=False, + ): + dpg.add_table_column(width=28, init_width_or_weight=28) + dpg.add_table_column(width=28, init_width_or_weight=28) + dpg.add_table_column(width=28, init_width_or_weight=28) + dpg.add_table_column(width=28, init_width_or_weight=28) + dpg.add_table_column(width=28, init_width_or_weight=28) + dpg.add_table_column(width=50, init_width_or_weight=50) + dpg.add_table_column(width=220, init_width_or_weight=220) + dpg.add_table_column(width=65, init_width_or_weight=65) + dpg.add_table_column(width=50, init_width_or_weight=50) + dpg.add_table_column(width=50, init_width_or_weight=50) - dpg.add_spacer(width=8) + with dpg.table_row(): + for star in range(1, 6): + filled = star <= entry.rating + label = "*" if filled else "o" + btn_tag = f"star_{entry.id}_{star}" + with dpg.table_cell(): + dpg.add_button( + label=label, + tag=btn_tag, + callback=_rate_by_id, + user_data=(entry.id, star), + width=24, + height=24, + ) - rating_text = ( - f"{entry.rating}\u2605" if entry.rating > 0 else "unrated" - ) - dpg.add_text(rating_text, tag=f"rating_label_{entry.id}") + with dpg.table_cell(): + rating_text = ( + f"{entry.rating}*" if entry.rating > 0 else "unrated" + ) + dpg.add_text(rating_text, tag=f"rating_label_{entry.id}") - dpg.add_spacer(width=12) + with dpg.table_cell(): + dpg.add_input_text( + hint="feedback...", + default_value=entry.feedback or "", + tag=f"feedback_{entry.id}", + width=200, + height=24, + ) - dpg.add_input_text( - hint="feedback...", - default_value=entry.feedback or "", - tag=f"feedback_{entry.id}", - width=200, - height=24, - ) + with dpg.table_cell(): + dpg.add_button( + label="Save", + callback=_save_feedback_by_id, + user_data=entry.id, + width=55, + height=24, + ) - dpg.add_spacer(width=4) + with dpg.table_cell(): + dpg.add_button( + label="Del", + callback=_delete_by_id, + user_data=entry.id, + width=40, + height=24, + ) - dpg.add_button( - label="Save", - callback=lambda s, a, e=entry: _save_feedback(e), - width=55, - height=24, - ) + with dpg.table_cell(): + dpg.add_button( + label="Params", + callback=_toggle_params, + user_data=entry.id, + width=50, + height=24, + ) - dpg.add_spacer(width=4) - - dpg.add_button( - label="Del", - callback=lambda s, a, e=entry: _delete(e), - width=40, - height=24, + params_group = f"params_{entry.id}" + with dpg.group(tag=params_group, show=False): + dpg.add_text( + _format_params(entry.params), + color=(160, 160, 160), + wrap=650, ) if entry.feedback: @@ -169,19 +207,6 @@ def _build_entry_row(entry): dpg.add_spacer(height=4) -def _rate(entry, star_tag): - star_num = int(star_tag.split("_")[-1]) - current = _get_current_rating(entry.id) - new_rating = 0 if current == star_num else star_num - - fb_widget = f"feedback_{entry.id}" - feedback = dpg.get_value(fb_widget) if dpg.does_item_exist(fb_widget) else "" - feedback = feedback if feedback else None - - _db.update_rating(entry.id, new_rating, feedback) - _refresh() - - def _get_current_rating(entry_id): for e in _entries: if e.id == entry_id: @@ -189,19 +214,51 @@ def _get_current_rating(entry_id): return 0 -def _save_feedback(entry): - fb_tag = f"feedback_{entry.id}" +def _rate_by_id(sender, app_data, user_data): + entry_id, star_num = user_data + current = _get_current_rating(entry_id) + new_rating = 0 if current == star_num else star_num + + fb_widget = f"feedback_{entry_id}" + feedback = dpg.get_value(fb_widget) if dpg.does_item_exist(fb_widget) else "" + feedback = feedback if feedback else None + + _db.update_rating(entry_id, new_rating, feedback) + _refresh() + + +def _save_feedback_by_id(sender, app_data, user_data): + entry_id = user_data + fb_tag = f"feedback_{entry_id}" feedback = dpg.get_value(fb_tag) if dpg.does_item_exist(fb_tag) else "" - rating = _get_current_rating(entry.id) - _db.update_rating(entry.id, rating, feedback if feedback else None) + rating = _get_current_rating(entry_id) + _db.update_rating(entry_id, rating, feedback if feedback else None) _refresh() -def _delete(entry): - _db.delete(entry.id) +def _delete_by_id(sender, app_data, user_data): + entry_id = user_data + _db.delete(entry_id) _refresh() +def _toggle_params(sender, app_data, user_data): + entry_id = user_data + params_group = f"params_{entry_id}" + if dpg.does_item_exist(params_group): + current = dpg.is_item_visible(params_group) + dpg.configure_item(params_group, show=not current) + + +def _format_params(params: dict) -> str: + lines = [] + for key, value in params.items(): + if key in ("model", "lora", "lcm") and isinstance(value, str): + value = os.path.basename(value) + lines.append(f"{key}: {value}") + return "\n".join(lines) + + def _on_filter(sender, app_data): global _filter _filter = dpg.get_value("filter_radio") diff --git a/server.py b/server.py index 1e67ab7..399c06a 100644 --- a/server.py +++ b/server.py @@ -282,6 +282,15 @@ def generate_sprite( "steps": steps, "remove_bg": remove_bg, "pixel_size": pixel_size, + "model": MODEL_DIR, + "lora": LORA_DIR, + "lcm": LCM_LORA_DIR, + "lora_scale": PIXEL_LORA_SCALE, + "lcm_scale": LCM_LORA_SCALE, + "negative_prompt": NEGATIVE_PROMPT, + "guidance_scale": 1.5, + "scheduler": "LCMScheduler", + "full_prompt": full_prompt, }, image_path=output_path, ) @@ -361,6 +370,15 @@ def batch_generate( "steps": steps, "remove_bg": remove_bg, "pixel_size": pixel_size, + "model": MODEL_DIR, + "lora": LORA_DIR, + "lcm": LCM_LORA_DIR, + "lora_scale": PIXEL_LORA_SCALE, + "lcm_scale": LCM_LORA_SCALE, + "negative_prompt": NEGATIVE_PROMPT, + "guidance_scale": 1.5, + "scheduler": "LCMScheduler", + "full_prompt": full_prompt, }, image_path=output_path, )