Views
No views yet
1from operator import itemgetter
2
3from detikzify.model import load
4from detikzify.infer import DetikzifyPipeline
5
6image = "https://w.wiki/A7Cc"
7pipeline = DetikzifyPipeline(*load(
8 model_name_or_path="nllg/detikzify-ds-7b",
9 device_map="auto",
10 torch_dtype="bfloat16",
11))
12
13# generate a single TikZ program
14fig = pipeline.sample(image=image)
15
16# if it compiles, rasterize it and show it
17if fig.is_rasterizable:
18 fig.rasterize().show()
19
20# run MCTS for 10 minutes and generate multiple TikZ programs
21figs = set()
22for score, fig in pipeline.simulate(image=image, timeout=600):
23 figs.add((score, fig))
24
25# save the best TikZ program
26best = sorted(figs, key=itemgetter(0))[-1][1]
27best.save("fig.tex")