Views
No views yet
CODIGPT2 (extends GPT2LMHeadModel)<|latent|>, <|start-latent|>, <|end-latent|> for latent reasoning<|start-latent|> tokenhuggingface-cli download dd101bb/latent-tts-codi --local-dir checkpoints/codi1from transformers import AutoTokenizer
2from src.generation_mixin import LatentGenerationMixin, LatentGenerationConfig
3from src.paths import MODELS
4
5# Load tokenizer
6model_id = "checkpoints/codi"
7tokenizer = AutoTokenizer.from_pretrained(model_id)
8if tokenizer.pad_token is None:
9 tokenizer.pad_token = tokenizer.eos_token
10
11# Get latent token IDs
12latent_id = tokenizer.convert_tokens_to_ids("<|latent|>")
13start_id = tokenizer.convert_tokens_to_ids("<|start-latent|>")
14end_id = tokenizer.convert_tokens_to_ids("<|end-latent|>")
15
16# Create model class with generation mixin
17class LatentCODI(MODELS["codi"]["class"], LatentGenerationMixin):
18 def __init__(self, config):
19 super().__init__(config)
20
21# Load model
22model = LatentCODI.from_pretrained(
23 model_id,
24 latent_id=latent_id,
25 latent_start_id=start_id,
26 latent_end_id=end_id,
27 device_map="auto",
28)
29
30# Prepare input (note: no newline before <|start-latent|>)
31question = "What is 2 + 2?<|start-latent|>"
32inputs = tokenizer(question, return_tensors="pt").to(model.device)
33
34# Configure generation
35generation_config = LatentGenerationConfig(
36 max_new_tokens=512,
37 latent_length=6,
38 latent_do_sample=True,
39 latent_do_sample_by="dropout", # or "noise"
40 dropout_p=0.1,
41 pad_token_id=tokenizer.pad_token_id,
42 eos_token_id=tokenizer.eos_token_id,
43)
44
45# Generate
46output = model.generate(
47 **inputs,
48 generation_config=generation_config,
49 num_return_sequences=1,
50)
51
52# Decode result
53result = tokenizer.decode(output[0], skip_special_tokens=True)
54print(result)1# Prepare batch inputs
2questions = [
3 "What is 2 + 2?<|start-latent|>",
4 "What is 5 * 3?<|start-latent|>",
5 "What is 10 - 4?<|start-latent|>",
6]
7inputs = tokenizer(questions, return_tensors="pt", padding=True).to(model.device)
8
9# Generate for batch
10outputs = model.generate(
11 **inputs,
12 generation_config=generation_config,
13 num_return_sequences=1,
14)
15
16# Decode batch results
17results = tokenizer.batch_decode(outputs, skip_special_tokens=True)
18for result in results:
19 print(result)1# Projector configuration (if enabled in model)
2projector = nn.Sequential(
3 nn.Dropout(projector_dropout),
4 nn.Linear(hidden_size, projector_hidden_size),
5 nn.GELU(),
6 nn.Linear(projector_hidden_size, hidden_size),
7 nn.LayerNorm(hidden_size),
8)output_hidden_states=True and config.projector=True.max_new_tokens (int): Maximum number of tokens to generatelatent_length (int): Number of latent tokens (default: 6)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 extract_answer_number
2
3# Extract answer from generated text
4answer = extract_answer_number(result)
5print(f"Answer: {answer}")1# For CODI (GPT-2 based models)
2./run_tests.sh1@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{shen2025codicompressingchainofthoughtcontinuous,
12 title={CODI: Compressing Chain-of-Thought into Continuous Space via Self-Distillation},
13 author={Zhenyi Shen and Hanqi Yan and Linhai Zhang and Zhanghao Hu and Yali Du and Yulan He},
14 year={2025},
15 eprint={2502.21074},
16 archivePrefix={arXiv},
17 primaryClass={cs.CL},
18 url={https://arxiv.org/abs/2502.21074},
19}