初始化项目,由ModelHub XC社区提供模型

Model: anujjamwal/OpenMath-Nemotron-1.5B-PruneAware
Source: Original Platform
This commit is contained in:
ModelHub XC
2026-07-19 14:49:13 +08:00
commit 0352e5741d
16 changed files with 913 additions and 0 deletions

36
.gitattributes vendored Normal file
View File

@@ -0,0 +1,36 @@
*.7z filter=lfs diff=lfs merge=lfs -text
*.arrow filter=lfs diff=lfs merge=lfs -text
*.bin filter=lfs diff=lfs merge=lfs -text
*.bz2 filter=lfs diff=lfs merge=lfs -text
*.ckpt filter=lfs diff=lfs merge=lfs -text
*.ftz filter=lfs diff=lfs merge=lfs -text
*.gz filter=lfs diff=lfs merge=lfs -text
*.h5 filter=lfs diff=lfs merge=lfs -text
*.joblib filter=lfs diff=lfs merge=lfs -text
*.lfs.* filter=lfs diff=lfs merge=lfs -text
*.mlmodel filter=lfs diff=lfs merge=lfs -text
*.model filter=lfs diff=lfs merge=lfs -text
*.msgpack filter=lfs diff=lfs merge=lfs -text
*.npy filter=lfs diff=lfs merge=lfs -text
*.npz filter=lfs diff=lfs merge=lfs -text
*.onnx filter=lfs diff=lfs merge=lfs -text
*.ot filter=lfs diff=lfs merge=lfs -text
*.parquet filter=lfs diff=lfs merge=lfs -text
*.pb filter=lfs diff=lfs merge=lfs -text
*.pickle filter=lfs diff=lfs merge=lfs -text
*.pkl filter=lfs diff=lfs merge=lfs -text
*.pt filter=lfs diff=lfs merge=lfs -text
*.pth filter=lfs diff=lfs merge=lfs -text
*.rar filter=lfs diff=lfs merge=lfs -text
*.safetensors filter=lfs diff=lfs merge=lfs -text
saved_model/**/* filter=lfs diff=lfs merge=lfs -text
*.tar.* filter=lfs diff=lfs merge=lfs -text
*.tar filter=lfs diff=lfs merge=lfs -text
*.tflite filter=lfs diff=lfs merge=lfs -text
*.tgz filter=lfs diff=lfs merge=lfs -text
*.wasm filter=lfs diff=lfs merge=lfs -text
*.xz filter=lfs diff=lfs merge=lfs -text
*.zip filter=lfs diff=lfs merge=lfs -text
*.zst filter=lfs diff=lfs merge=lfs -text
*tfevents* filter=lfs diff=lfs merge=lfs -text
tokenizer.json filter=lfs diff=lfs merge=lfs -text

67
README.md Normal file
View File

@@ -0,0 +1,67 @@
---
base_model: anujjamwal/OpenMath-Nemotron-1.5B-PruneAware
library_name: transformers
model_name: OpenMath-Nemotron-1.5B-PruneAware
tags:
- generated_from_trainer
- sft
- trl
- custom_generate
licence: license
datasets:
- anujjamwal/OpenMathReasoning-Sampled-Hierarchical-Cot
---
# Model Card for OpenMath-Nemotron-1.5B-PruneAware
This model implements [Cognitive Compression](https://github.com/anujjamwal/cognitive-compression) an approach to produce hierarchical
structured chain of thought that can be actively pruned at inference time while maintaining the solution quality.
Tradition Chain-of-Thought is append-onl; a token once generated remains in context for ever. Context compression introduces hierarchical
reasoning where reasoning is broken into subproblems. Once the subproblem is solved, its full chain of thought can be discarded and
replaced with the **summary and solution** dramatically reducing the context window pressure.
This model is a fine-tuned version of [anujjamwal/OpenMath-Nemotron-1.5B-PruneAware](https://huggingface.co/anujjamwal/OpenMath-Nemotron-1.5B-PruneAware).
It has been trained using [TRL](https://github.com/huggingface/trl).
## Quick start
```python
from transformers import pipeline
question = "If you had a time machine, but could only go to the past or the future once and never return, which would you choose and why?"
generator = pipeline("text-generation", model="anujjamwal/OpenMath-Nemotron-1.5B-PruneAware", device="cuda")
output = generator([{"role": "user", "content": question}], max_new_tokens=128, return_full_text=False)[0]
print(output["generated_text"])
```
## Training procedure
This model was trained with SFT.
### Framework versions
- TRL: 0.29.0
- Transformers: 5.0.0
- Pytorch: 2.10.0+cu128
- Datasets: 4.0.0
- Tokenizers: 0.22.2
## Citations
Cite TRL as:
```bibtex
@misc{jamwal2026cognitivecompression,
title = {{Cognitive Compression: Hierarchical Chain of Thought for Efficient LLM Reasoning}},
author = {Jamwal, Anuj},
url = {huggingface.co/anujjamwal/OpenMath-Nemotron-1.5B-PruneAware},
year = {2026},
note = {CS224N Winter '26 Final Project: Stanford University}
}
```

21
chat_template.jinja Normal file
View File

@@ -0,0 +1,21 @@
{%- if messages[0]['role'] == 'system' -%}
<|im_start|>system
{{ messages[0]['content'] | trim }}<|im_end|>
{%- else -%}
<|im_start|>system
<|im_end|>
{%- endif -%}
{%- for message in messages -%}
{%- if (message.role == 'user') or (message.role == 'system' and not loop.first) or (message.role == 'assistant') -%}
<|im_start|>{{ message.role }}
{%- if message['role'] == 'assistant' %}
{% generation %}{{ message['content'] | trim }}<|im_end|>{% endgeneration %}
{%- else %}
{{ message['content'] | trim }}<|im_end|>
{%- endif %}
{%- endif -%}
{%- endfor -%}
{%- if add_generation_prompt %}
<|im_start|>assistant
{%- endif -%}

61
config.json Normal file
View File

@@ -0,0 +1,61 @@
{
"architectures": [
"Qwen2ForCausalLM"
],
"attention_dropout": 0.0,
"bos_token_id": null,
"dtype": "bfloat16",
"eos_token_id": 151645,
"hidden_act": "silu",
"hidden_size": 1536,
"initializer_range": 0.02,
"intermediate_size": 8960,
"layer_types": [
"full_attention",
"full_attention",
"full_attention",
"full_attention",
"full_attention",
"full_attention",
"full_attention",
"full_attention",
"full_attention",
"full_attention",
"full_attention",
"full_attention",
"full_attention",
"full_attention",
"full_attention",
"full_attention",
"full_attention",
"full_attention",
"full_attention",
"full_attention",
"full_attention",
"full_attention",
"full_attention",
"full_attention",
"full_attention",
"full_attention",
"full_attention",
"full_attention"
],
"max_position_embeddings": 131072,
"max_window_layers": 21,
"model_type": "qwen2",
"num_attention_heads": 12,
"num_hidden_layers": 28,
"num_key_value_heads": 2,
"pad_token_id": 151643,
"rms_norm_eps": 1e-06,
"rope_parameters": {
"rope_theta": 500000.0,
"rope_type": "default"
},
"sliding_window": null,
"tie_word_embeddings": true,
"transformers_version": "5.0.0",
"use_cache": false,
"use_sliding_window": false,
"vocab_size": 151670
}

670
custom_generate/generate.py Normal file
View File

@@ -0,0 +1,670 @@
from typing import Any, Optional, Sequence, Tuple
import torch
from torch import nn
from trl.trainer.utils import get_config_model_id
from transformers import AutoProcessor, Cache, DynamicCache, LogitsProcessorList, ProcessorMixin, StoppingCriteriaList
from transformers.generation.utils import GenerationMixin, ALL_CACHE_NAMES, GenerateEncoderDecoderOutput, GenerateDecoderOnlyOutput
from transformers.generation.configuration_utils import GenerationConfig
from transformers.generation.streamers import BaseStreamer
from transformers.utils.generic import ModelOutput
from transformers import PreTrainedTokenizerBase
# ---------------------------------------------------------------------------
# KV-cache manipulation
# ---------------------------------------------------------------------------
def _retain_and_prune_kv_cache(
cache: DynamicCache,
prune_map: dict[int, Tuple[int, int]],
batch_size: int,
old_seq_len: int,
) -> int:
"""Modify a DynamicCache in-place after a pruning event.
For each pruned batch element the function keeps only the prefix
``[0..thought_pos]`` in the cache. Both thought and solution tokens are
removed. The caller sets ``cache_position`` so the next forward pass
re-processes the solution tokens against the cached prefix.
For non-pruned batch elements all entries are kept.
Compatible with both transformers ≥5.x (``cache.layers[i].keys/values``)
and older versions (``cache.key_cache[i]``).
Returns the new cache sequence length (``max_new_seq``).
"""
# Detect API version: transformers ≥5 uses cache.layers, older uses cache.key_cache
use_layers_api = hasattr(cache, 'layers')
if use_layers_api:
num_layers = len(cache.layers)
elif hasattr(cache, 'key_cache'):
num_layers = len(cache.key_cache) # type: ignore[attr-defined]
else:
return 0
if num_layers == 0:
return 0
if use_layers_api:
first_keys = cache.layers[0].keys if len(cache.layers) > 0 else None
else:
first_keys = cache.key_cache[0] if len(cache.key_cache) > 0 else None # type: ignore[attr-defined]
if first_keys is None:
return 0
device = first_keys.device
dtype = first_keys.dtype
head_dim = first_keys.shape[-1]
# Compute per-element prefix length as plain ints (no GPU tensors needed).
# The kept positions are always contiguous [0..N) by construction, so we
# only need the length, not explicit index tensors.
prefix_lengths: list[int] = []
for b in range(batch_size):
if b in prune_map:
thought_pos, _solution_pos = prune_map[b]
prefix_lengths.append(thought_pos + 1)
else:
prefix_lengths.append(old_seq_len)
max_new_seq = max(prefix_lengths)
# Fast path: all batch elements keep the same contiguous prefix length.
# Common case (batch_size=1, or all pruned to the same point).
all_same_prefix = len(set(prefix_lengths)) == 1
# Pre-compute gather indices once (used across all layers) for the
# heterogeneous fallback path only.
if not all_same_prefix:
idx = torch.zeros(batch_size, max_new_seq, dtype=torch.long, device=device)
for b in range(batch_size):
n = prefix_lengths[b]
idx[b, :n] = torch.arange(n, device=device)
for layer_idx in range(num_layers):
if use_layers_api:
old_keys = cache.layers[layer_idx].keys # (batch, heads, seq, dim)
old_vals = cache.layers[layer_idx].values
else:
old_keys = cache.key_cache[layer_idx] # type: ignore[attr-defined]
old_vals = cache.value_cache[layer_idx] # type: ignore[attr-defined]
if all_same_prefix:
# Single slice — no allocation, just a contiguous copy
new_keys = old_keys[:, :, :prefix_lengths[0], :].contiguous()
new_vals = old_vals[:, :, :prefix_lengths[0], :].contiguous()
else:
# Vectorized gather across batch and heads
num_heads = old_keys.shape[1]
idx_expanded = idx[:, None, :, None].expand(-1, num_heads, -1, head_dim)
new_keys = torch.gather(old_keys, 2, idx_expanded)
new_vals = torch.gather(old_vals, 2, idx_expanded)
# Zero out padding positions beyond each element's prefix.
# Without this, gather copies position-0 values into padding slots,
# which FA2 would attend to (it ignores attention masks).
for b in range(batch_size):
n = prefix_lengths[b]
if n < max_new_seq:
new_keys[b, :, n:, :] = 0
new_vals[b, :, n:, :] = 0
if use_layers_api:
cache.layers[layer_idx].keys = new_keys
cache.layers[layer_idx].values = new_vals
else:
cache.key_cache[layer_idx] = new_keys # type: ignore[attr-defined]
cache.value_cache[layer_idx] = new_vals # type: ignore[attr-defined]
if hasattr(cache, '_seen_tokens'):
cache._seen_tokens = max_new_seq
return max_new_seq
# ---------------------------------------------------------------------------
# Generation helpers
# ---------------------------------------------------------------------------
def _prepare_inputs_for_generation(
model,
input_ids: torch.LongTensor,
past_key_values: Cache | None = None,
attention_mask: torch.LongTensor | None = None,
inputs_embeds: torch.FloatTensor | None = None,
cache_position: torch.LongTensor | None = None,
is_first_iteration: bool | None = False,
**kwargs,
):
# After a prune event, cache_position covers only the solution tokens
# that need re-processing (fewer than input_ids). The standard
# prepare_inputs_for_generation expects them to match, so pre-slice
# input_ids to the tokens at cache_position. This also handles the
# normal decode case where cache_position is a single element.
if (
cache_position is not None
and past_key_values is not None
and input_ids.shape[1] != cache_position.shape[0]
):
input_ids = input_ids[:, cache_position]
model_inputs = GenerationMixin.prepare_inputs_for_generation(
model,
input_ids,
past_key_values=past_key_values,
attention_mask=attention_mask,
inputs_embeds=inputs_embeds,
cache_position=cache_position,
is_first_iteration=is_first_iteration,
**kwargs
)
return model_inputs
def _update_model_kwargs_for_generation(
model,
outputs: ModelOutput,
model_kwargs: dict[str, Any],
is_encoder_decoder: bool = False,
num_new_tokens: int = 1,
use_pos_ids_buf: bool = False,
) -> dict[str, Any]:
model_kwargs = GenerationMixin._update_model_kwargs_for_generation(
model,
outputs,
model_kwargs,
is_encoder_decoder=is_encoder_decoder,
num_new_tokens=num_new_tokens
)
# HF's _update concatenates new positions onto cache_position, which is
# correct for the initial prefill→decode transition but breaks after a
# prune: the multi-element cache_position from _prune_model_inputs gets
# carried forward, causing prepare_inputs_for_generation to re-select
# already-cached tokens on every subsequent step. DynamicCache.update()
# appends these duplicates, inflating the cache until old_seq_len exceeds
# the pruned input length and torch.arange crashes. Fix: always keep
# only the last element (the next decode position).
if "cache_position" in model_kwargs and model_kwargs["cache_position"] is not None:
model_kwargs["cache_position"] = model_kwargs["cache_position"][-1:]
# HF's _update_model_kwargs_for_generation does not manage position_ids.
# For prune-agnostic mode we track position_ids explicitly after each
# prune event so that RoPE positions match training.
# When use_pos_ids_buf=True, _sample manages position_ids via a
# pre-allocated buffer, so skip the O(n) concatenation here.
if (
"position_ids" in model_kwargs
and model_kwargs["position_ids"] is not None
and not use_pos_ids_buf
):
pos = model_kwargs["position_ids"]
model_kwargs["position_ids"] = torch.cat([pos, pos[:, -1:] + 1], dim=-1)
return model_kwargs
def _prune_model_inputs(
model,
prune_input_candidates: Sequence[int],
prune_input_locations: Sequence[Sequence[Tuple[int, int, int]]],
input_ids: torch.LongTensor,
prune_aware: bool,
model_kwargs: dict[str, Any],
retain_kv_cache: bool = True,
) -> Tuple[torch.LongTensor, dict[str, Any]]:
"""Prune input sequences after a ``[RETURN]`` token is generated.
When ``retain_kv_cache=True`` (default), the function retains the
prefix (tokens before the thought block) in the KV cache and sets
``cache_position`` so the next forward pass re-processes the solution
tokens against the cached prefix. Cost: O(k) where k = number of
solution tokens.
When ``retain_kv_cache=False``, the entire cache is discarded and a
full re-prefill is performed. Cost: O(N).
"""
is_prune_agnostic = not prune_aware
batch_size = input_ids.shape[0]
device = input_ids.device
if is_prune_agnostic:
# Construct position_ids from model_kwargs if already tracked,
# otherwise build contiguous positions (correct before first prune).
if "position_ids" in model_kwargs and model_kwargs["position_ids"] is not None:
position_ids = model_kwargs["position_ids"]
else:
position_ids = torch.arange(
input_ids.shape[1], device=device
).unsqueeze(0).expand(batch_size, -1)
else:
position_ids = None
# Map batch index → (thought_pos, solution_pos) for valid prune targets
prune_map: dict[int, Tuple[int, int]] = {}
for cand_idx, batch_idx in enumerate(prune_input_candidates):
for thought_pos, solution_pos, _return_pos in prune_input_locations[cand_idx]:
if solution_pos is not None:
prune_map[batch_idx] = (thought_pos, solution_pos)
if not prune_map:
return input_ids, model_kwargs
# ------------------------------------------------------------------
# Decide whether to retain the KV cache
# ------------------------------------------------------------------
cache = model_kwargs.get('past_key_values', None)
use_layers_api = hasattr(cache, 'layers')
if cache is None:
cache_populated = False
elif use_layers_api:
cache_populated = len(cache.layers) > 0
elif hasattr(cache, 'key_cache'):
cache_populated = len(cache.key_cache) > 0 # type: ignore[union-attr]
else:
cache_populated = False
can_retain_cache = (
retain_kv_cache
and cache is not None
and isinstance(cache, DynamicCache)
and cache_populated
)
# ------------------------------------------------------------------
# Build pruned rows per batch element (same as before)
# ------------------------------------------------------------------
new_rows: list[torch.Tensor] = []
if is_prune_agnostic:
new_position_rows: list[torch.Tensor] = list(
position_ids[b] for b in range(batch_size) # type: ignore
)
for b in range(batch_size):
if b in prune_map:
thought_pos, solution_pos = prune_map[b]
new_rows.append(torch.cat((input_ids[b, :thought_pos + 1], input_ids[b, solution_pos + 1:])))
if is_prune_agnostic:
new_position_rows[b] = torch.cat((position_ids[b, :thought_pos + 1], position_ids[b, solution_pos + 1:])) # type: ignore
else:
new_rows.append(input_ids[b])
# Build pruned input_ids, attention_mask, and position_ids
if batch_size == 1:
# Fast path: no inter-batch padding needed
new_input_ids = new_rows[0].unsqueeze(0)
max_len = new_input_ids.shape[1]
model_kwargs['attention_mask'] = torch.ones(1, max_len, dtype=torch.long, device=device)
if is_prune_agnostic:
model_kwargs['position_ids'] = new_position_rows[0].unsqueeze(0)
else:
model_kwargs['position_ids'] = None
else:
# Pad to uniform length across batch
max_len = max(r.shape[0] for r in new_rows)
pad_id = getattr(model.config, 'pad_token_id', None)
if pad_id is None:
pad_id = getattr(model.config, 'eos_token_id', 0)
new_input_ids = torch.full((batch_size, max_len), pad_id, dtype=input_ids.dtype, device=device)
new_attention_mask = torch.zeros((batch_size, max_len), dtype=torch.long, device=device)
for b, r in enumerate(new_rows):
new_input_ids[b, :r.shape[0]] = r
new_attention_mask[b, :r.shape[0]] = 1
if is_prune_agnostic:
new_position_ids = torch.zeros((batch_size, max_len), dtype=torch.long, device=device)
for b, p in enumerate(new_position_rows):
new_position_ids[b, :p.shape[0]] = p
model_kwargs['position_ids'] = new_position_ids
else:
model_kwargs['position_ids'] = None
model_kwargs['attention_mask'] = new_attention_mask
# ------------------------------------------------------------------
# KV cache handling
# ------------------------------------------------------------------
if can_retain_cache:
if hasattr(cache, 'layers') and len(cache.layers) > 0:
old_seq_len = cache.layers[0].keys.shape[2]
else:
old_seq_len = cache.key_cache[0].shape[2] # type: ignore[union-attr]
new_cache_seq = _retain_and_prune_kv_cache(
cache=cache, # type: ignore[arg-type]
prune_map=prune_map,
batch_size=batch_size,
old_seq_len=old_seq_len,
)
# Cache has only the prefix (a tokens). Set cache_position so
# the next forward pass re-processes the solution tokens (c tokens)
# against the cached prefix. Cost: O(k) where k = solution length.
model_kwargs['cache_position'] = torch.arange(
new_cache_seq, max_len, dtype=torch.int64, device=device,
)
else:
# Fallback: discard the KV cache and re-prefill from scratch
old_cache = model_kwargs.pop('past_key_values', None)
del old_cache
# Reset cache_position to cover the full pruned sequence so that
# prepare_inputs_for_generation treats the next forward as a prefill
model_kwargs['cache_position'] = torch.arange(max_len, dtype=torch.int64, device=device)
if position_ids is not None:
del position_ids
return new_input_ids, model_kwargs # type: ignore
def _sample(
model,
input_ids: torch.LongTensor,
logits_processor: LogitsProcessorList,
stopping_criteria: StoppingCriteriaList,
generation_config: GenerationConfig,
processing_class: Optional[PreTrainedTokenizerBase] = None,
synced_gpus: bool = False,
streamer: Optional["BaseStreamer"] = None,
prune_aware: bool = False,
retain_kv_cache: bool = True,
return_unpruned_output: bool = False,
**model_kwargs,
):
"""Generate sequences using argmax or sampling from model logits.
This function implements the core generation loop for token-by-token sequence generation.
It supports both deterministic (argmax) and stochastic (sampling) token selection, and includes
special handling for pruning sequences based on custom tokens ([THOUGHT], [SOLUTION], [RETURN]).
Args:
model: The language model used for generation.
input_ids (torch.LongTensor): Initial input token IDs of shape (batch_size, seq_len).
logits_processor (LogitsProcessorList): List of processors to apply to logits before sampling.
stopping_criteria (StoppingCriteriaList): List of criteria to determine when to stop generation.
generation_config (GenerationConfig): Configuration object containing generation parameters.
processing_class (Optional[PreTrainedTokenizerBase]): Tokenizer for token-id conversion. Defaults to None.
synced_gpus (bool): Whether GPUs are synchronized for multi-device generation. Defaults to False.
streamer (Optional[BaseStreamer]): Optional streamer to output tokens during generation. Defaults to None.
prune_aware (bool): When True, positions are renumbered after pruning (contiguous). Defaults to False.
retain_kv_cache (bool): When True (default), retains the prefix in the KV cache
after pruning and re-processes only the solution tokens. Cost: O(k).
When False, discards the cache entirely and re-prefills. Cost: O(N).
**model_kwargs: Additional keyword arguments passed to the model forward pass.
Returns:
torch.LongTensor: Generated sequences of shape (batch_size, seq_len) including input and generated tokens.
Notes:
- With retain_kv_cache=True, the prefix is cached and solution tokens are
re-processed via a standard forward pass.
- With retain_kv_cache=False, the KV cache is discarded and the full
pruned sequence is re-prefilled from scratch.
- In prune-agnostic mode, position_ids are tracked explicitly to maintain original RoPE
positions for surviving tokens.
- Handles batch generation with unfinished sequences tracking.
- Manages KV-cache through model_kwargs for efficient decoding.
"""
if processing_class is None:
processing_class = AutoProcessor.from_pretrained(
get_config_model_id(model.config), truncation_side="left", padding_side="left"
)
# Handle pad token for processors or tokenizers
if isinstance(processing_class, ProcessorMixin):
tokenizer = processing_class.tokenizer
elif isinstance(processing_class, PreTrainedTokenizerBase):
tokenizer = processing_class
else:
raise TypeError("The `processing_class` must be either a `PreTrainedTokenizerBase` or a `ProcessorMixin`")
if tokenizer.pad_token is None:
tokenizer.pad_token = tokenizer.eos_token
thought_token_id = processing_class.convert_tokens_to_ids("[THOUGHT]")
solution_token_id = processing_class.convert_tokens_to_ids("[SOLUTION]")
return_token_id = processing_class.convert_tokens_to_ids("[RETURN]")
pad_token_id = generation_config._pad_token_tensor # type: ignore
output_attentions = generation_config.output_attentions
output_hidden_states = generation_config.output_hidden_states
output_scores = generation_config.output_scores
output_logits = generation_config.output_logits
return_dict_in_generate = generation_config.return_dict_in_generate
has_eos_stopping_criteria = any(hasattr(criteria, "eos_token_id") for criteria in stopping_criteria)
do_sample = generation_config.do_sample
scores = () if (return_dict_in_generate and output_scores) else None
raw_logits = () if (return_dict_in_generate and output_logits) else None
decoder_attentions = () if (return_dict_in_generate and output_attentions) else None
cross_attentions = () if (return_dict_in_generate and output_attentions) else None
decoder_hidden_states = () if (return_dict_in_generate and output_hidden_states) else None
batch_size = input_ids.shape[0]
this_peer_finished: bool = False
unfinished_sequences = torch.ones(batch_size, dtype=torch.long, device=input_ids.device)
stacks = [[] for _ in range(batch_size)]
if return_unpruned_output:
unpruned_ids = [input_ids[b].tolist() for b in range(batch_size)]
model_forward = (
model.get_compiled_call(generation_config.compile_config)
if GenerationMixin._valid_auto_compile_criteria(model, model_kwargs, generation_config)
else model.__call__
)
# Pre-allocate input_ids buffer to avoid O(n²) copies from torch.cat
# on every token step. input_ids becomes a view into this buffer.
_pad_id_scalar = pad_token_id.item() if isinstance(pad_token_id, torch.Tensor) else pad_token_id
_max_new = generation_config.max_new_tokens if generation_config.max_new_tokens is not None else generation_config.max_length
_buf_len = input_ids.shape[1] + _max_new
_ids_buf = torch.full(
(batch_size, _buf_len), _pad_id_scalar,
dtype=input_ids.dtype, device=input_ids.device,
)
_ids_buf[:, :input_ids.shape[1]] = input_ids
_cur_len = input_ids.shape[1]
input_ids = _ids_buf[:, :_cur_len]
# Pre-allocate position_ids buffer for prune-agnostic mode to avoid
# O(n²) torch.cat growth in _update_model_kwargs_for_generation.
_pos_ids_buf = None
if not prune_aware:
_pos_ids_buf = torch.zeros(
(batch_size, _buf_len), dtype=torch.long, device=input_ids.device,
)
# Assisted generation completes the prefill stage in candidate generator so that
# we don't have several `prefill` calls in one generation loop. Skip `_prefill` for assistants
if not generation_config.is_assistant:
outputs = GenerationMixin._prefill(model, input_ids, generation_config, model_kwargs)
prefill_consumed = False
else:
model_kwargs = GenerationMixin._get_initial_cache_position(model, input_ids.shape[1], input_ids.device, model_kwargs)
prefill_consumed = True
while GenerationMixin._has_unfinished_sequences(model, this_peer_finished, synced_gpus, device=input_ids.device):
if prefill_consumed:
model_inputs = _prepare_inputs_for_generation(model, input_ids, **model_kwargs)
with GenerationMixin._optimize_model_for_decode(model):
outputs = model_forward(**model_inputs, return_dict=True)
prefill_consumed = True
model_kwargs = _update_model_kwargs_for_generation(
model,
outputs, # type: ignore The value is always initialized
model_kwargs,
is_encoder_decoder=model.config.is_encoder_decoder,
use_pos_ids_buf=(_pos_ids_buf is not None),
)
# Extend position_ids via buffer instead of torch.cat (prune-agnostic mode)
if _pos_ids_buf is not None and model_kwargs.get("position_ids") is not None:
pos_ids = model_kwargs["position_ids"]
plen = pos_ids.shape[1]
_pos_ids_buf[:, :plen] = pos_ids
_pos_ids_buf[:, plen] = _pos_ids_buf[:, plen - 1] + 1
model_kwargs["position_ids"] = _pos_ids_buf[:, :plen + 1]
if synced_gpus and this_peer_finished:
continue
# .float() severs the reference to outputs.logits (avoids keeping the
# full logits tensor alive) and converts to float32 in a single op.
# The redundant device= kwarg is removed (logits are already on the
# correct device). The del outputs below provides additional safety.
next_token_logits = outputs.logits[:, -1, :].float() # type: ignore
# pre-process distribution
next_token_scores = logits_processor(input_ids, next_token_logits)
# token selection
if do_sample:
probs = nn.functional.softmax(next_token_scores, dim=-1)
next_tokens = torch.multinomial(probs, num_samples=1).squeeze(1)
else:
next_tokens = torch.argmax(next_token_scores, dim=-1)
# finished sentences should have their next token be a padding token
if has_eos_stopping_criteria:
next_tokens = next_tokens * unfinished_sequences + pad_token_id * (1 - unfinished_sequences)
# update generated ids, model inputs, and length for next step
# Write into pre-allocated buffer instead of O(n²) torch.cat
_ids_buf[:, _cur_len] = next_tokens
_cur_len += 1
input_ids = _ids_buf[:, :_cur_len]
# Single GPU→CPU transfer; all comparisons on CPU to avoid
# multiple synchronization stalls per token step.
# Also used by streamer and return_unpruned_output.
next_tokens_cpu = next_tokens.cpu()
if streamer is not None:
streamer.put(next_tokens_cpu)
pos = input_ids.shape[1] - 1
if return_unpruned_output:
for b in range(batch_size):
unpruned_ids[b].append(next_tokens_cpu[b].item())
for b in (next_tokens_cpu == thought_token_id).nonzero(as_tuple=True)[0].tolist():
stacks[b].append([pos, None, None])
for b in (next_tokens_cpu == solution_token_id).nonzero(as_tuple=True)[0].tolist():
if stacks[b]:
stacks[b][-1][1] = pos
is_return = next_tokens_cpu == return_token_id
for b in is_return.nonzero(as_tuple=True)[0].tolist():
if stacks[b]:
stacks[b][-1][2] = pos
return_indices = is_return.nonzero(as_tuple=True)[0].tolist()
prune_candidates = [idx for idx in return_indices if stacks[idx]]
if prune_candidates:
input_ids, model_kwargs = _prune_model_inputs(
model,
prune_input_candidates=prune_candidates,
prune_input_locations=[[stacks[b].pop()] for b in prune_candidates],
input_ids=input_ids,
prune_aware=prune_aware,
model_kwargs=model_kwargs,
retain_kv_cache=retain_kv_cache,
)
# Copy pruned result back into the pre-allocated buffer.
# No need to zero the tail — it was initialized to pad_id and
# input_ids is always sliced to _cur_len so stale data is never read.
_cur_len = input_ids.shape[1]
_ids_buf[:, :_cur_len] = input_ids
input_ids = _ids_buf[:, :_cur_len]
# Sync position_ids into buffer after pruning (prune-agnostic mode)
if _pos_ids_buf is not None and model_kwargs.get("position_ids") is not None:
pos_ids = model_kwargs["position_ids"]
_pos_ids_buf[:, :pos_ids.shape[1]] = pos_ids
model_kwargs["position_ids"] = _pos_ids_buf[:, :pos_ids.shape[1]]
unfinished_sequences = unfinished_sequences & ~stopping_criteria(input_ids, scores) # type: ignore
this_peer_finished = not unfinished_sequences.any()
# This is needed to properly delete outputs.logits which may be very large for first iteration
# Otherwise a reference to outputs is kept which keeps the logits alive in the next iteration
del outputs # type: ignore
if streamer is not None:
streamer.end()
if return_unpruned_output:
max_len = max(len(ids) for ids in unpruned_ids)
pad_id = pad_token_id.item() if isinstance(pad_token_id, torch.Tensor) else pad_token_id
unpruned_tensor = torch.full(
(batch_size, max_len), pad_id,
dtype=input_ids.dtype, device=input_ids.device,
)
for b, ids in enumerate(unpruned_ids):
unpruned_tensor[b, :len(ids)] = torch.tensor(ids, dtype=input_ids.dtype, device=input_ids.device)
input_ids = unpruned_tensor
if return_dict_in_generate:
cache = None
if any(cache_key in model_kwargs for cache_key in ALL_CACHE_NAMES):
cache_key = next(cache_key for cache_key in ALL_CACHE_NAMES if cache_key in model_kwargs)
cache = model_kwargs[cache_key]
if model.config.is_encoder_decoder:
return GenerateEncoderDecoderOutput(
sequences=input_ids,
scores=scores,
logits=raw_logits,
encoder_attentions=encoder_attentions,
encoder_hidden_states=encoder_hidden_states,
decoder_attentions=decoder_attentions,
cross_attentions=cross_attentions,
decoder_hidden_states=decoder_hidden_states,
past_key_values=cache,
)
else:
return GenerateDecoderOnlyOutput(
sequences=input_ids,
scores=scores,
logits=raw_logits,
attentions=decoder_attentions,
hidden_states=decoder_hidden_states,
past_key_values=cache,
)
else:
return input_ids
def generate(
model,
processing_class = None,
retain_kv_cache: bool = True,
return_unpruned_output: bool = False,
prune_aware: bool = True,
**kwargs
):
"""Custom generate method for Hierarchical Chain of Thought.
Args:
model: The language model.
processing_class: Tokenizer / processing class.
retain_kv_cache: When True (default), retains the prefix in the KV
cache after pruning and re-processes solution tokens. When False,
discards the cache entirely.
**kwargs: Forwarded to ``model.generate``.
"""
custom_generate = kwargs.pop('custom_generate', _sample)
return GenerationMixin.generate(
model,
custom_generate=custom_generate,
processing_class=processing_class,
retain_kv_cache=retain_kv_cache,
return_unpruned_output=return_unpruned_output,
prune_aware=prune_aware,
**kwargs
)

8
generation_config.json Normal file
View File

@@ -0,0 +1,8 @@
{
"_from_model_config": true,
"eos_token_id": [
151645
],
"pad_token_id": 151643,
"transformers_version": "5.0.0"
}

3
model.safetensors Normal file
View File

@@ -0,0 +1,3 @@
version https://git-lfs.github.com/spec/v1
oid sha256:7f874629d9e4ae8fe591c35b90d9be466ffe46b14517c4374921bb24c7d2dc27
size 3086649992

View File

@@ -0,0 +1,3 @@
version https://git-lfs.github.com/spec/v1
oid sha256:1a8aa4dfcdfe2a62ea4dc50004ab26f2dc6ae539a0eb2285d48a224ca9bb80b9
size 4184

View File

@@ -0,0 +1,3 @@
version https://git-lfs.github.com/spec/v1
oid sha256:732159757ddc9ed12837baa05d0389ecffb5c9c8c6cad209e07d60efb07f1875
size 4184

View File

@@ -0,0 +1,3 @@
version https://git-lfs.github.com/spec/v1
oid sha256:6b7511eac4998041fc1f88e2a732f32ccd51c57fccb8a4c4d118f997bee9de4d
size 86660

View File

@@ -0,0 +1,3 @@
version https://git-lfs.github.com/spec/v1
oid sha256:7f218382ba36c6f18f9a9a600e152ee523632722acb24a7e7e78cfc3aa9a0483
size 6377

View File

@@ -0,0 +1,3 @@
version https://git-lfs.github.com/spec/v1
oid sha256:572491d57b30e4aaf2aae18956aa26456966fd81223dff84a8f72dd0034bbaa3
size 18046

View File

@@ -0,0 +1,3 @@
version https://git-lfs.github.com/spec/v1
oid sha256:1e942e470b55b5e34d73ded4e49fec144c2f873729d35c06ce5a7e1573fd4bf2
size 37057

3
tokenizer.json Normal file
View File

@@ -0,0 +1,3 @@
version https://git-lfs.github.com/spec/v1
oid sha256:4aa0f6873d276e78d29759ff1ba764fe84358db27b39fe8501119662f4c1c55a
size 11422819

23
tokenizer_config.json Normal file
View File

@@ -0,0 +1,23 @@
{
"add_prefix_space": false,
"backend": "tokenizers",
"bos_token": null,
"clean_up_tokenization_spaces": false,
"eos_token": "<|im_end|>",
"errors": "replace",
"extra_special_tokens": [
"<|im_end|>",
"<|endoftext|>",
"[THOUGHT]",
"[SOLUTION]",
"[RETURN]",
"<think>",
"</think>"
],
"is_local": false,
"model_max_length": 131072,
"pad_token": "<|endoftext|>",
"split_special_tokens": false,
"tokenizer_class": "Qwen2Tokenizer",
"unk_token": null
}

3
training_args.bin Normal file
View File

@@ -0,0 +1,3 @@
version https://git-lfs.github.com/spec/v1
oid sha256:e1f17aa1966f73b2b5887f58e68fe868b017bc34bc0e5330299055e295525502
size 5713