commit 0352e5741d06101bc48ca2fb24a08dfe3534ed3f Author: ModelHub XC Date: Sun Jul 19 14:49:13 2026 +0800 初始化项目,由ModelHub XC社区提供模型 Model: anujjamwal/OpenMath-Nemotron-1.5B-PruneAware Source: Original Platform diff --git a/.gitattributes b/.gitattributes new file mode 100644 index 0000000..52373fe --- /dev/null +++ b/.gitattributes @@ -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 diff --git a/README.md b/README.md new file mode 100644 index 0000000..1577f26 --- /dev/null +++ b/README.md @@ -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} +} +``` \ No newline at end of file diff --git a/chat_template.jinja b/chat_template.jinja new file mode 100644 index 0000000..2cfe475 --- /dev/null +++ b/chat_template.jinja @@ -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 -%} + \ No newline at end of file diff --git a/config.json b/config.json new file mode 100644 index 0000000..882eca2 --- /dev/null +++ b/config.json @@ -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 +} diff --git a/custom_generate/generate.py b/custom_generate/generate.py new file mode 100644 index 0000000..5378baa --- /dev/null +++ b/custom_generate/generate.py @@ -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 + ) diff --git a/generation_config.json b/generation_config.json new file mode 100644 index 0000000..d81914d --- /dev/null +++ b/generation_config.json @@ -0,0 +1,8 @@ +{ + "_from_model_config": true, + "eos_token_id": [ + 151645 + ], + "pad_token_id": 151643, + "transformers_version": "5.0.0" +} diff --git a/model.safetensors b/model.safetensors new file mode 100644 index 0000000..990f783 --- /dev/null +++ b/model.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:7f874629d9e4ae8fe591c35b90d9be466ffe46b14517c4374921bb24c7d2dc27 +size 3086649992 diff --git a/runs/Mar05_05-44-50_43bf35ae8d9d/events.out.tfevents.1772689490.43bf35ae8d9d.826.0 b/runs/Mar05_05-44-50_43bf35ae8d9d/events.out.tfevents.1772689490.43bf35ae8d9d.826.0 new file mode 100644 index 0000000..106297b --- /dev/null +++ b/runs/Mar05_05-44-50_43bf35ae8d9d/events.out.tfevents.1772689490.43bf35ae8d9d.826.0 @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:1a8aa4dfcdfe2a62ea4dc50004ab26f2dc6ae539a0eb2285d48a224ca9bb80b9 +size 4184 diff --git a/runs/Mar05_05-45-13_43bf35ae8d9d/events.out.tfevents.1772689513.43bf35ae8d9d.826.1 b/runs/Mar05_05-45-13_43bf35ae8d9d/events.out.tfevents.1772689513.43bf35ae8d9d.826.1 new file mode 100644 index 0000000..941ff1f --- /dev/null +++ b/runs/Mar05_05-45-13_43bf35ae8d9d/events.out.tfevents.1772689513.43bf35ae8d9d.826.1 @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:732159757ddc9ed12837baa05d0389ecffb5c9c8c6cad209e07d60efb07f1875 +size 4184 diff --git a/runs/Mar05_05-47-42_43bf35ae8d9d/events.out.tfevents.1772689662.43bf35ae8d9d.3741.0 b/runs/Mar05_05-47-42_43bf35ae8d9d/events.out.tfevents.1772689662.43bf35ae8d9d.3741.0 new file mode 100644 index 0000000..0b9f75c --- /dev/null +++ b/runs/Mar05_05-47-42_43bf35ae8d9d/events.out.tfevents.1772689662.43bf35ae8d9d.3741.0 @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:6b7511eac4998041fc1f88e2a732f32ccd51c57fccb8a4c4d118f997bee9de4d +size 86660 diff --git a/runs/Mar05_14-53-41_13b8607adfc1/events.out.tfevents.1772722421.13b8607adfc1.1928.0 b/runs/Mar05_14-53-41_13b8607adfc1/events.out.tfevents.1772722421.13b8607adfc1.1928.0 new file mode 100644 index 0000000..247071e --- /dev/null +++ b/runs/Mar05_14-53-41_13b8607adfc1/events.out.tfevents.1772722421.13b8607adfc1.1928.0 @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:7f218382ba36c6f18f9a9a600e152ee523632722acb24a7e7e78cfc3aa9a0483 +size 6377 diff --git a/runs/Mar05_14-56-44_13b8607adfc1/events.out.tfevents.1772722604.13b8607adfc1.1928.1 b/runs/Mar05_14-56-44_13b8607adfc1/events.out.tfevents.1772722604.13b8607adfc1.1928.1 new file mode 100644 index 0000000..df1f6ec --- /dev/null +++ b/runs/Mar05_14-56-44_13b8607adfc1/events.out.tfevents.1772722604.13b8607adfc1.1928.1 @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:572491d57b30e4aaf2aae18956aa26456966fd81223dff84a8f72dd0034bbaa3 +size 18046 diff --git a/runs/Mar05_15-28-41_13b8607adfc1/events.out.tfevents.1772724521.13b8607adfc1.1928.2 b/runs/Mar05_15-28-41_13b8607adfc1/events.out.tfevents.1772724521.13b8607adfc1.1928.2 new file mode 100644 index 0000000..8db658b --- /dev/null +++ b/runs/Mar05_15-28-41_13b8607adfc1/events.out.tfevents.1772724521.13b8607adfc1.1928.2 @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:1e942e470b55b5e34d73ded4e49fec144c2f873729d35c06ce5a7e1573fd4bf2 +size 37057 diff --git a/tokenizer.json b/tokenizer.json new file mode 100644 index 0000000..e85ad4d --- /dev/null +++ b/tokenizer.json @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:4aa0f6873d276e78d29759ff1ba764fe84358db27b39fe8501119662f4c1c55a +size 11422819 diff --git a/tokenizer_config.json b/tokenizer_config.json new file mode 100644 index 0000000..ad5085e --- /dev/null +++ b/tokenizer_config.json @@ -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]", + "", + "" + ], + "is_local": false, + "model_max_length": 131072, + "pad_token": "<|endoftext|>", + "split_special_tokens": false, + "tokenizer_class": "Qwen2Tokenizer", + "unk_token": null +} diff --git a/training_args.bin b/training_args.bin new file mode 100644 index 0000000..f846521 --- /dev/null +++ b/training_args.bin @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:e1f17aa1966f73b2b5887f58e68fe868b017bc34bc0e5330299055e295525502 +size 5713