mirror of
https://github.com/NVIDIA/Model-Optimizer.git
synced 2026-10-02 03:14:52 +08:00
## What does this PR do? **Type of change:** new feature **Overview:** Training and inference code for Dynamic Memory Sparsification (DMS) - method from NeurIPS 2025 paper [Inference-Time Hyper-Scaling with KV Cache Compression](https://neurips.cc/virtual/2025/loc/san-diego/poster/119605) ## Usage Detailed in `experimental/dms/README.md` and `experimental/dms/ARCHITECTURE.md` ## Testing DMS tests in `experimental/dms/tests` covering: * prefill * generation * gradient propagation * chunked prefill ## Before your PR is "*Ready for review*" <!-- If you haven't finished some of the above items you can still open `Draft` PR. --> - **Make sure you read and follow [Contributor guidelines](https://github.com/NVIDIA/Model-Optimizer/blob/main/CONTRIBUTING.md)** and your commits are signed. - **Is this change backward compatible?**: Yes - **Did you write any new necessary tests?**: Yes - **Did you add or update any necessary documentation?**: Yes - **Did you update [Changelog](https://github.com/NVIDIA/Model-Optimizer/blob/main/CHANGELOG.rst)?**: No, DMS is currently experimental feature with description in `experimental/dms` ## Additional Information A minimal, optimized implementation of the DMS algorithm for KV-cache compression, as described in: > **Inference-Time Hyper-Scaling with KV Cache Compression** > Adrian Łańcucki, Konrad Staniszewski, Piotr Nawrot, Edoardo M. Ponti > Paper: [https://arxiv.org/abs/2506.05345](https://arxiv.org/abs/2506.05345) > NeurIPS: [https://neurips.cc/virtual/2025/loc/san-diego/poster/119605](https://neurips.cc/virtual/2025/loc/san-diego/poster/119605) Inference-time scaling trades efficiency for improved reasoning by generating longer sequences. In Transformer LLMs, generation cost is often bottlenecked by the size of the key-value (KV) cache. DMS addresses this by learning a KV cache eviction policy that compresses the cache while preserving accuracy. ## How it works DMS learns a per-head eviction policy that determines which KV cache entries to keep during generation. Rather than immediately discarding tokens, DMS delays eviction decisions, implicitly merging representations and preserving critical information. During training, the compression ratio is gradually increased from 1× to a target value (e.g., 8×), using knowledge distillation to match the outputs of an uncompressed teacher model. <!-- This is an auto-generated comment: release notes by coderabbit.ai --> ## Summary by CodeRabbit * **New Features** * Introduces Dynamic Memory Sparsification (DMS), an algorithm for efficient LLM inference and training with adaptive attention gating. * Adds DMS-enabled Qwen3 models with memory-efficient KV cache management and paged block-based storage. * Includes student-teacher distillation training infrastructure with noise scheduling and compression ratio control. * Provides configuration system and training/evaluation scripts for DMS adaptation. * **Documentation** * Added architecture guide, README, and example inference notebook. * **Tests** * Added comprehensive test suite for chunked prefill, cache management, and prefill/inference validation. <!-- end of auto-generated comment: release notes by coderabbit.ai --> --------- Signed-off-by: Konrad Staniszewski <kstaniszewsk@nvidia.com> Signed-off-by: kstaniszewsknv <kstaniszewsk@nvidia.com> Co-authored-by: coderabbitai[bot] <136622811+coderabbitai[bot]@users.noreply.github.com>
33 lines
1.0 KiB
Python
33 lines
1.0 KiB
Python
# SPDX-FileCopyrightText: Copyright (c) 2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
|
# SPDX-License-Identifier: Apache-2.0
|
|
#
|
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
|
# you may not use this file except in compliance with the License.
|
|
# You may obtain a copy of the License at
|
|
#
|
|
# http://www.apache.org/licenses/LICENSE-2.0
|
|
#
|
|
# Unless required by applicable law or agreed to in writing, software
|
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
|
# See the License for the specific language governing permissions and
|
|
# limitations under the License.
|
|
|
|
|
|
"""Utility functions for DMS tests."""
|
|
|
|
import os
|
|
import sys
|
|
|
|
import pytest
|
|
|
|
|
|
def add_dms_to_path():
|
|
"""Add the DMS package to the Python path."""
|
|
sys.path.append(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
|
|
|
|
|
ignore_flex_attention_warnings = pytest.mark.filterwarnings(
|
|
"ignore:flex_attention called without torch.compile:UserWarning",
|
|
)
|