文生TikZ

目录

Qwen-GeoGebra-Coder-7B

文生TikZ

Qwen/Qwen3-0.6B 要单独下载


Qwen-GeoGebra-Coder-7B

from fastapi import FastAPI
from fastapi.middleware.cors import CORSMiddleware
from llama_cpp import Llama
import uvicorn
import re

app = FastAPI()

# Standard CORS setup for local web tools
app.add_middleware(
    CORSMiddleware,
    allow_origins=["*"],
    allow_methods=["*"],
    allow_headers=["*"]
)

# Load the fine-tuned math model
# Utilizing the full 16GB VRAM of the RTX 4060 Ti
llm = Llama(
    model_path="/home/khurram/ai_models/math_dataset/math_viz_Q4_K_M.gguf",
    n_gpu_layers=-1, 
    n_ctx=2048,
    n_batch=512,
    temperature=0.1
)

def clean_and_format_ggb(raw_text):
    """
    Standardizes model coordinates and brackets for GeoGebra.
    Converts [Cylinder[<0,0,0>, <3,0,0>, <0,10,0>]] 
    to Cylinder((0,0,0), (0,10,0), 3.0)
    """
    # 1. Clean up bracket variations
    text = raw_text.replace("<", "(").replace(">", ")")
    
    # 2. Extract all coordinate sets (x,y,z)
    coords = re.findall(r"\((-?\d+\.?\d*,\s*-?\d+\.?\d*,\s*-?\d+\.?\d*)\)", text)
    
    if len(coords) >= 3:
        bottom_pt = f"({coords[0]})"
        top_pt = f"({coords[2]})"
        
        # Extract scalar radius from the middle point
        radius_match = re.findall(r"[-+]?\d*\.\d+|\d+", coords[1])
        radius = next((abs(float(n)) for n in radius_match if float(n) != 0), 3.0)
        
        return f"Cylinder({bottom_pt}, {top_pt}, {radius})"
    
    # Fallback for simple Sphere or direct commands
    return text.replace("[", "").replace("]", "").replace("<", "(").replace(">", ")").strip()

@app.post("/ask")
async def ask_geo(data: dict):
    user_prompt = data.get("prompt", "")
    
    # Step 1: Request the "Thought" (Mathematical Reasoning)
    prompt = f"<|im_start|>user\n{user_prompt}<|im_end|>\n<|im_start|>thought\n"
    thought_output = llm(prompt, max_tokens=150, stop=["<|im_end|>"])
    thought_text = thought_output["choices"][0]["text"].strip()

    # Step 2: Request the "Assistant" (GeoGebra Code)
    command_prompt = f"{prompt}{thought_text}<|im_end|>\n<|im_start|>assistant\n"
    command_output = llm(command_prompt, max_tokens=150, stop=["<|im_end|>"])
    assistant_raw = command_output["choices"][0]["text"].strip()

    # Final string formatting for the GeoGebra Applet
    final_cmds = clean_and_format_ggb(assistant_raw)
    
    return {
        "commands": final_cmds, 
        "thought": thought_text
    }

if __name__ == "__main__":
    uvicorn.run(app, host="127.0.0.1", port=8000)

文生TikZ

环境安装:

pip install peft

pip install transformers==5.14.1

Qwen/Qwen3-0.6B 要单独下载

import torch
from transformers import AutoModelForCausalLM, AutoTokenizer
from peft import PeftModel

BASE = "Qwen/Qwen3-0.6B"                 # match the adapter's base model
REPO = "kyhe/qwen3-geotikz"              # this Hub repo
ADAPTER = "qwen3-pgf-geotikz"            # subfolder for the v2 specialist

tok = AutoTokenizer.from_pretrained(BASE)
model = AutoModelForCausalLM.from_pretrained(BASE, dtype=torch.bfloat16)
model = PeftModel.from_pretrained(model, REPO, subfolder=ADAPTER).eval()

SYSTEM = (
    "You are a geometry-to-TikZ compiler. Given a geometry scene described only through "
    "relationships and constraints (no explicit coordinates), you must derive the exact "
    "coordinates yourself and output a single valid TikZ/PGF figure that compiles and "
    "renders the described geometry. Output ONLY the TikZ code, starting with "
    "\\begin{tikzpicture} and ending with \\end{tikzpicture}. No prose, no markdown fences."
)
scene = ("There is a circle centered at the origin with radius 3. Point A lies on the "
         "circle at 40 degrees. Point B lies on the circle at 200 degrees. Point M is the "
         "midpoint of segment AB.")
msgs = [{"role": "system", "content": SYSTEM},
        {"role": "user", "content": f"Scene:\n{scene}\n\nReturn the TikZ figure."}]
prompt = tok.apply_chat_template(msgs, tokenize=False, add_generation_prompt=True,
                                 enable_thinking=False)
out = model.generate(**tok(prompt, return_tensors="pt").to(model.device),
                     max_new_tokens=512, do_sample=False)
print(tok.decode(out[0][tok(prompt, return_tensors="pt")["input_ids"].shape[1]:],
                 skip_special_tokens=True))

cannot import name 'EncoderDecoderCache' from 'transformers'

评论
添加红包

请填写红包祝福语或标题

红包个数最小为10个

红包金额最低5元

当前余额3.43前往充值 >
需支付:10.00
成就一亿技术人!
领取后你会自动成为博主和红包主的粉丝 规则
hope_wisdom
发出的红包

打赏作者

AI算法网奇

你的鼓励将是我创作的最大动力

¥1 ¥2 ¥4 ¥6 ¥10 ¥20
扫码支付:¥1
获取中
扫码支付

您的余额不足,请更换扫码支付或充值

打赏作者

实付
使用余额支付
点击重新获取
扫码支付
钱包余额 0

抵扣说明:

1.余额是钱包充值的虚拟货币,按照1:1的比例进行支付金额的抵扣。
2.余额无法直接购买下载,可以购买VIP、付费专栏及课程。

余额充值