Add generation params to DB, fix GUI star layout/refresh, add sample generator script

This commit is contained in:
Emil
2026-06-28 16:56:44 +03:00
parent 68cba874d5
commit 09f6f377c7
4 changed files with 223 additions and 60 deletions
+3
View File
@@ -164,6 +164,9 @@ class FeedbackDB:
"instruction": f"Generate a pixel-art sprite for: {entry.prompt}", "instruction": f"Generate a pixel-art sprite for: {entry.prompt}",
"response": entry.params, "response": entry.params,
"rating": entry.rating, "rating": entry.rating,
"feedback": entry.feedback,
"image_path": entry.image_path,
"prompt": entry.prompt,
} }
) )
lines.append(line) lines.append(line)
+85
View File
@@ -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())
+117 -60
View File
@@ -28,6 +28,8 @@ _texture_registry = None
def _load_thumbnail(path, tag): def _load_thumbnail(path, tag):
if not path or not os.path.exists(path): if not path or not os.path.exists(path):
return False return False
if dpg.does_item_exist(tag):
return True
try: try:
img = Image.open(path).convert("RGBA") img = Image.open(path).convert("RGBA")
img.thumbnail((THUMB_SIZE, THUMB_SIZE), Image.NEAREST) 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_text(p, color=(100, 100, 100))
dpg.add_spacer(height=4) dpg.add_spacer(height=4)
with dpg.group(horizontal=True): with dpg.table(
for star in range(1, 6): header_row=False,
filled = star <= entry.rating policy=dpg.mvTable_SizingFixedFit,
label = "\u2605" if filled else "\u2606" no_host_extendX=True,
color = (255, 200, 80) if filled else (80, 80, 90) row_background=False,
btn_tag = f"star_{entry.id}_{star}" borders_innerH=False,
dpg.add_button( borders_outerH=False,
label=label, borders_innerV=False,
tag=btn_tag, borders_outerV=False,
callback=lambda s, a, e=entry: _rate(e, s), ):
width=28, dpg.add_table_column(width=28, init_width_or_weight=28)
height=24, dpg.add_table_column(width=28, init_width_or_weight=28)
) dpg.add_table_column(width=28, init_width_or_weight=28)
dpg.bind_item_theme(btn_tag, _star_theme(color)) 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 = ( with dpg.table_cell():
f"{entry.rating}\u2605" if entry.rating > 0 else "unrated" rating_text = (
) f"{entry.rating}*" if entry.rating > 0 else "unrated"
dpg.add_text(rating_text, tag=f"rating_label_{entry.id}") )
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( with dpg.table_cell():
hint="feedback...", dpg.add_button(
default_value=entry.feedback or "", label="Save",
tag=f"feedback_{entry.id}", callback=_save_feedback_by_id,
width=200, user_data=entry.id,
height=24, 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( with dpg.table_cell():
label="Save", dpg.add_button(
callback=lambda s, a, e=entry: _save_feedback(e), label="Params",
width=55, callback=_toggle_params,
height=24, user_data=entry.id,
) width=50,
height=24,
)
dpg.add_spacer(width=4) params_group = f"params_{entry.id}"
with dpg.group(tag=params_group, show=False):
dpg.add_button( dpg.add_text(
label="Del", _format_params(entry.params),
callback=lambda s, a, e=entry: _delete(e), color=(160, 160, 160),
width=40, wrap=650,
height=24,
) )
if entry.feedback: if entry.feedback:
@@ -169,19 +207,6 @@ def _build_entry_row(entry):
dpg.add_spacer(height=4) 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): def _get_current_rating(entry_id):
for e in _entries: for e in _entries:
if e.id == entry_id: if e.id == entry_id:
@@ -189,19 +214,51 @@ def _get_current_rating(entry_id):
return 0 return 0
def _save_feedback(entry): def _rate_by_id(sender, app_data, user_data):
fb_tag = f"feedback_{entry.id}" 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 "" feedback = dpg.get_value(fb_tag) if dpg.does_item_exist(fb_tag) else ""
rating = _get_current_rating(entry.id) rating = _get_current_rating(entry_id)
_db.update_rating(entry.id, rating, feedback if feedback else None) _db.update_rating(entry_id, rating, feedback if feedback else None)
_refresh() _refresh()
def _delete(entry): def _delete_by_id(sender, app_data, user_data):
_db.delete(entry.id) entry_id = user_data
_db.delete(entry_id)
_refresh() _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): def _on_filter(sender, app_data):
global _filter global _filter
_filter = dpg.get_value("filter_radio") _filter = dpg.get_value("filter_radio")
+18
View File
@@ -282,6 +282,15 @@ def generate_sprite(
"steps": steps, "steps": steps,
"remove_bg": remove_bg, "remove_bg": remove_bg,
"pixel_size": pixel_size, "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, image_path=output_path,
) )
@@ -361,6 +370,15 @@ def batch_generate(
"steps": steps, "steps": steps,
"remove_bg": remove_bg, "remove_bg": remove_bg,
"pixel_size": pixel_size, "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, image_path=output_path,
) )