Buckets:
6.81 GB
14 files
Updated 5 months ago
Ctrl+K
| Name | Size | Uploaded | Xet hash |
|---|---|---|---|
| .gitattributes | 101 Bytes xet | a7e3ea59 | |
| README.md | 2.51 kB xet | d1004b6c | |
| added_tokens.json | 605 Bytes xet | 652df9a4 | |
| chat_template.jinja | 2.43 kB xet | 8a2b5062 | |
| config.json | 1.52 kB xet | 30e6b965 | |
| generation_config.json | 117 Bytes xet | c6780c8a | |
| merges.txt | 1.67 MB xet | 87912eed | |
| model-00001-of-00002.safetensors | 4.96 GB xet | c8c7b62c | |
| model-00002-of-00002.safetensors | 1.84 GB xet | 36672b33 | |
| model.safetensors.index.json | 35.7 kB xet | 302670dc | |
| special_tokens_map.json | 616 Bytes xet | c176b813 | |
| tokenizer.json | 11.4 MB xet | ff206f79 | |
| tokenizer_config.json | 4.69 kB xet | 7efa4f58 | |
| vocab.json | 2.78 MB xet | 9208e1be |
Model Card for TikZilla-3B-RL
TikZilla-3B-RL is a language model for generating TikZ/LaTeX figures from natural language descriptions.
It is based on TikZilla-3B and was trained with reinforcement learning (RL) on DaTikZ-V4 for scientific figure generation.
Installation
pip install torch==2.5.1 transformers==4.53.2 accelerate==1.8.1
Usage
import torch
from transformers import AutoTokenizer, AutoModelForCausalLM, GenerationConfig
model_id = "nllg/TikZilla-3B-RL"
tokenizer = AutoTokenizer.from_pretrained(model_id)
model = AutoModelForCausalLM.from_pretrained(
model_id,
torch_dtype=torch.bfloat16,
device_map="auto",
)
eos_token_id = tokenizer.convert_tokens_to_ids("<|im_end|>")
pad_token_id = tokenizer.pad_token_id or tokenizer.eos_token_id or eos_token_id
gen_config = GenerationConfig(
do_sample=True,
temperature=1.0,
top_p=0.9,
max_new_tokens=2048,
eos_token_id=eos_token_id,
pad_token_id=pad_token_id,
)
your_input_description = "A scientific line plot showing two curves. The x-axis is labeled 'Time' ranging from 0 to 100, and the y-axis is labeled 'Value' ranging from 0 to 1. The first curve is a blue solid line that gradually increases from near 0 and levels off around 0.9. The second curve is a red dashed line that rises to a peak around the middle of the plot and then decreases. A legend in the upper right labels the blue line as 'Model A' and the red dashed line as 'Model B'. The background is white with light gray grid lines."
messages = [
{
"role": "user",
"content": (
"Generate a complete LaTeX document that contains a TikZ figure according to the following requirements:\n"
+ your_input_description +
"\nWrap your code using \\documentclass[tikz]{standalone}, and include \\begin{document}...\\end{document}. "
"Only output valid LaTeX code with no extra text."
),
}
]
text = tokenizer.apply_chat_template(
messages,
tokenize=False,
add_generation_prompt=True,
)
inputs = tokenizer([text], return_tensors="pt").to(model.device)
output_ids = model.generate(**inputs, generation_config=gen_config)
response_ids = output_ids[0][len(inputs["input_ids"][0]):]
output = tokenizer.decode(response_ids, skip_special_tokens=True)
print(output)
- Total size
- 6.81 GB
- Files
- 14
- Last updated
- May 24
- Pre-warmed CDN
- US EU US EU