dcode / app.py
twarner's picture
Add model dropdown, improve UI
0073cbe
Raw
History Blame
4.8 kB
"""dcode Gradio Space - Text to Gcode inference."""
import re
import gradio as gr
import torch
from transformers import AutoModelForSeq2SeqLM, AutoModelForCausalLM, AutoTokenizer
# Available models
MODELS = {
"flan-t5-base (best)": "twarner/dcode-flan-t5-base",
}
# Machine limits
BOUNDS = {"left": -420.5, "right": 420.5, "top": 594.5, "bottom": -594.5}
# Cache loaded models
_model_cache = {}
def get_model(model_name: str):
"""Load and cache model."""
if model_name not in _model_cache:
model_id = MODELS[model_name]
device = "cuda" if torch.cuda.is_available() else "cpu"
dtype = torch.float16 if device == "cuda" else torch.float32
tokenizer = AutoTokenizer.from_pretrained(model_id)
if "gpt2" in model_id or "codegen" in model_id:
model = AutoModelForCausalLM.from_pretrained(model_id, torch_dtype=dtype).to(device)
else:
model = AutoModelForSeq2SeqLM.from_pretrained(model_id, torch_dtype=dtype).to(device)
model.eval()
_model_cache[model_name] = (model, tokenizer, device)
return _model_cache[model_name]
def validate_gcode(gcode: str) -> str:
"""Clamp coordinates to machine bounds."""
lines = []
for line in gcode.split("\n"):
corrected = line
x_match = re.search(r"X([-\d.]+)", line, re.IGNORECASE)
if x_match:
x = float(x_match.group(1))
x = max(BOUNDS["left"], min(BOUNDS["right"], x))
corrected = re.sub(r"X[-\d.]+", f"X{x:.2f}", corrected, flags=re.IGNORECASE)
y_match = re.search(r"Y([-\d.]+)", line, re.IGNORECASE)
if y_match:
y = float(y_match.group(1))
y = max(BOUNDS["bottom"], min(BOUNDS["top"], y))
corrected = re.sub(r"Y[-\d.]+", f"Y{y:.2f}", corrected, flags=re.IGNORECASE)
lines.append(corrected)
return "\n".join(lines)
def generate(prompt: str, model_name: str, temperature: float, max_tokens: int) -> str:
"""Generate gcode from prompt."""
if not prompt or not prompt.strip():
return "Enter a prompt to generate gcode"
try:
model, tokenizer, device = get_model(model_name)
model_id = MODELS[model_name]
inputs = tokenizer(prompt, return_tensors="pt", max_length=128, truncation=True)
inputs = {k: v.to(device) for k, v in inputs.items()}
with torch.no_grad():
outputs = model.generate(
**inputs,
max_new_tokens=max_tokens,
do_sample=True,
temperature=temperature,
top_p=0.9,
pad_token_id=tokenizer.eos_token_id,
)
# For causal models, skip the input tokens
if "gpt2" in model_id or "codegen" in model_id:
gcode = tokenizer.decode(outputs[0][inputs["input_ids"].shape[1]:], skip_special_tokens=True)
else:
gcode = tokenizer.decode(outputs[0], skip_special_tokens=True)
gcode = validate_gcode(gcode)
line_count = len(gcode.split("\n"))
return f"; dcode output - {line_count} lines\n; Model: {model_name}\n; Machine validated\n\n{gcode}"
except Exception as e:
return f"; Error: {e}"
demo = gr.Interface(
fn=generate,
inputs=[
gr.Textbox(label="Prompt", placeholder="drawing of a cat...", lines=2),
gr.Dropdown(choices=list(MODELS.keys()), value="flan-t5-base (best)", label="Model"),
gr.Slider(0.1, 1.5, value=0.8, label="Temperature", info="Higher = more creative"),
gr.Slider(256, 2048, value=1024, step=256, label="Max Tokens"),
],
outputs=gr.Code(label="Gcode", language=None, lines=25),
title="dcode",
description="**Text → Polargraph Gcode** | Generate machine-compatible gcode from natural language. [GitHub](https://github.com/Twarner491/dcode) | [Model](https://huggingface.co/twarner/dcode-flan-t5-base) | [Dataset](https://huggingface.co/datasets/twarner/dcode-polargraph-gcode)",
examples=[
["drawing of a cat", "flan-t5-base (best)", 0.8, 1024],
["abstract spiral pattern", "flan-t5-base (best)", 0.9, 1024],
["simple house with chimney", "flan-t5-base (best)", 0.7, 512],
["portrait of a woman", "flan-t5-base (best)", 0.8, 1024],
],
theme=gr.themes.Soft(primary_hue="emerald"),
article="""
## About
dcode finetunes text-to-text models to directly output polargraph-compatible gcode from natural language descriptions.
**Training**: Flan-T5-base trained on 175,952 art-caption-gcode triplets for 20 epochs on H100.
**Machine Bounds**: X: ±420.5mm, Y: ±594.5mm | Pen servo: 40° (down) / 90° (up)
**License**: MIT
""",
)
if __name__ == "__main__":
demo.launch()