Views
No views yet
ColarLlama (extends LlamaForCausalLM)<|latent|> for latent reasoning### as the end-of-latent markerhuggingface-cli download dd101bb/latent-tts-colar --local-dir checkpoints/colar1import torch
2from transformers import AutoTokenizer
3from src.generation_mixin import LatentGenerationMixin, LatentGenerationConfig
4from src.paths import MODELS
5
6# Load tokenizer
7model_id = "checkpoints/colar"
8tokenizer = AutoTokenizer.from_pretrained(model_id)
9if tokenizer.pad_token is None:
10 tokenizer.pad_token = tokenizer.eos_token
11
12# Get latent token IDs
13latent_id = tokenizer.convert_tokens_to_ids("<|latent|>")
14end_id = tokenizer.convert_tokens_to_ids("###")
15
16# Create model class with generation mixin
17class LatentCoLaR(MODELS["colar"]["class"], LatentGenerationMixin):
18 pass
19
20# Load model
21model = LatentCoLaR.from_pretrained(
22 model_id,
23 device_map="auto",
24 torch_dtype=torch.bfloat16, # Recommended for LLaMA models
25)
26
27# Prepare input
28question = "What is 2 + 2?<|latent|>"
29inputs = tokenizer(question, return_tensors="pt").to(model.device)
30
31# Configure generation
32generation_config = LatentGenerationConfig(
33 max_new_tokens=128,
34 max_latent_length=64, # CoLaR uses max_latent_length instead of latent_length
35 latent_do_sample=True,
36 latent_do_sample_by="dropout", # or "noise"
37 dropout_p=0.1,
38 pad_token_id=tokenizer.pad_token_id,
39 eos_token_id=tokenizer.eos_token_id,
40)
41
42# Generate
43output = model.generate(
44 **inputs,
45 generation_config=generation_config,
46 num_return_sequences=1,
47)
48
49# Decode result
50result = tokenizer.decode(output[0], skip_special_tokens=True)
51print(result)1import torch
2
3# Prepare batch inputs
4questions = [
5 "What is 2 + 2?<|latent|>",
6 "What is 5 * 3?<|latent|>",
7 "What is 10 - 4?<|latent|>",
8]
9inputs = tokenizer(questions, return_tensors="pt", padding=True).to(model.device)
10
11# Generate for batch
12outputs = model.generate(
13 **inputs,
14 generation_config=generation_config,
15 num_return_sequences=1,
16)
17
18# Decode batch results
19results = tokenizer.batch_decode(outputs, skip_special_tokens=True)
20for result in results:
21 print(result)1class LatentHead(nn.Module):
2 def __init__(self, feature_size, intermediate_size=512):
3 super().__init__()
4 self.fc = nn.Sequential(
5 nn.Linear(feature_size, intermediate_size),
6 nn.GELU(),
7 nn.Linear(intermediate_size, intermediate_size),
8 nn.LayerNorm(intermediate_size),
9 )
10 self.mean = nn.Linear(intermediate_size, feature_size)latent_embedding_std (default: 0.018 for LLaMA-3.2 models).max_new_tokens (int): Maximum number of tokens to generatemax_latent_length (int): Maximum number of latent tokens (default: 64)latent_do_sample (bool): Whether to use stochastic samplinglatent_do_sample_by (str): Sampling method - "dropout" or "noise"dropout_p (float): Dropout probability for Monte Carlo Dropout (e.g., 0.1)noise_std (float): Standard deviation for Additive Gaussian Noise1generation_config = LatentGenerationConfig(
2 latent_do_sample_by="dropout",
3 dropout_p=0.1,
4 # ...
5)1generation_config = LatentGenerationConfig(
2 latent_do_sample_by="noise",
3 noise_std=0.1,
4 # ...
5)1from src.paths import colar_extract_answer_number
2
3# Extract answer from generated text
4answer = colar_extract_answer_number(result)
5print(f"Answer: {answer}")1# For CoLaR (LLaMA based models)
2./run_tests_llama.shtorch.bfloat16 or torch.float16 for LLaMA modelsmax_latent_length instead of fixed latent_length1@misc{you2025paralleltesttimescalinglatent,
2 title={Parallel Test-Time Scaling for Latent Reasoning Models},
3 author={Runyang You and Yongqi Li and Meng Liu and Wenjie Wang and Liqiang Nie and Wenjie Li},
4 year={2025},
5 eprint={2510.07745},
6 archivePrefix={arXiv},
7 primaryClass={cs.CL},
8 url={https://arxiv.org/abs/2510.07745},
9}
10
11@misc{tan2025thinksilentlythinkfast,
12 title={Think Silently, Think Fast: Dynamic Latent Compression of LLM Reasoning Chains},
13 author={Wenhui Tan and Jiaze Li and Jianzhong Ju and Zhenbo Luo and Jian Luan and Ruihua Song},
14 year={2025},
15 eprint={2505.16552},
16 archivePrefix={arXiv},
17 primaryClass={cs.CL},
18 url={https://arxiv.org/abs/2505.16552},
19}