初始化项目,由ModelHub XC社区提供模型
Model: anujjamwal/OpenMath-Nemotron-1.5B-PruneAware Source: Original Platform
This commit is contained in:
36
.gitattributes
vendored
Normal file
36
.gitattributes
vendored
Normal 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
67
README.md
Normal 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
21
chat_template.jinja
Normal 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
61
config.json
Normal 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
670
custom_generate/generate.py
Normal 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
8
generation_config.json
Normal 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
3
model.safetensors
Normal file
@@ -0,0 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:7f874629d9e4ae8fe591c35b90d9be466ffe46b14517c4374921bb24c7d2dc27
|
||||
size 3086649992
|
||||
@@ -0,0 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:1a8aa4dfcdfe2a62ea4dc50004ab26f2dc6ae539a0eb2285d48a224ca9bb80b9
|
||||
size 4184
|
||||
@@ -0,0 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:732159757ddc9ed12837baa05d0389ecffb5c9c8c6cad209e07d60efb07f1875
|
||||
size 4184
|
||||
@@ -0,0 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:6b7511eac4998041fc1f88e2a732f32ccd51c57fccb8a4c4d118f997bee9de4d
|
||||
size 86660
|
||||
@@ -0,0 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:7f218382ba36c6f18f9a9a600e152ee523632722acb24a7e7e78cfc3aa9a0483
|
||||
size 6377
|
||||
@@ -0,0 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:572491d57b30e4aaf2aae18956aa26456966fd81223dff84a8f72dd0034bbaa3
|
||||
size 18046
|
||||
@@ -0,0 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:1e942e470b55b5e34d73ded4e49fec144c2f873729d35c06ce5a7e1573fd4bf2
|
||||
size 37057
|
||||
3
tokenizer.json
Normal file
3
tokenizer.json
Normal file
@@ -0,0 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:4aa0f6873d276e78d29759ff1ba764fe84358db27b39fe8501119662f4c1c55a
|
||||
size 11422819
|
||||
23
tokenizer_config.json
Normal file
23
tokenizer_config.json
Normal 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
3
training_args.bin
Normal file
@@ -0,0 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:e1f17aa1966f73b2b5887f58e68fe868b017bc34bc0e5330299055e295525502
|
||||
size 5713
|
||||
Reference in New Issue
Block a user