Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
121 changes: 121 additions & 0 deletions modelopt/torch/models/gpt_oss/modeling_ptq.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,121 @@
# 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.

"""GPT-OSS PTQ modeling (HF model type ``gpt_oss``)."""

from contextlib import contextmanager
from functools import partial

import torch
from transformers.models.gpt_oss.modeling_gpt_oss import GptOssExperts

from modelopt.torch.quantization.nn import QuantModuleRegistry, TensorQuantizer
from modelopt.torch.quantization.plugins.custom import _QuantFunctionalMixin
from modelopt.torch.quantization.plugins.huggingface import (
_transposed_quantize,
_TransposedExpertsCalibMixin,
)

__all__: list[str] = []


class _QuantGptOssExperts(_TransposedExpertsCalibMixin, _QuantFunctionalMixin):
"""Quantized wrapper for `transformers.GptOssExperts`.

Quantizes `gate_up_proj` and `down_proj` weights via dynamic attributes inside `quantize_weight()`.
Activations into `gate_up_proj` are quantized by `gate_up_proj_input_quantizer`. For `down_proj`
activation quantization, we intercept `torch.Tensor.__matmul__`/`torch.bmm` and quantize inputs
on every second call (since the first call computes `gate_up_proj` outputs and second call
computes `down_proj` outputs).
"""

@staticmethod
def _get_quantized_weight(quantizer, module, weight):
# MoE weight is accessed for each expert in one forward pass. so lets cache it
if module._enable_weight_quantization:
if hasattr(quantizer, "_cached_quant_val"):
return getattr(quantizer, "_cached_quant_val")
quantizer._cached_quant_val = _transposed_quantize(weight, quantizer)
return quantizer._cached_quant_val
return weight

def _setup_for_weight_quantization(self):
self._register_dynamic_attribute(
"gate_up_proj", partial(self._get_quantized_weight, self.gate_up_proj_weight_quantizer)
)
self._register_dynamic_attribute(
"down_proj", partial(self._get_quantized_weight, self.down_proj_weight_quantizer)
)

def _setup(self):
assert not hasattr(self, "kernel_layer_name"), (
"ModelOpt quantization does not support patched forward for kernel_hub"
)
self.gate_up_proj_input_quantizer = TensorQuantizer()
self.gate_up_proj_weight_quantizer = TensorQuantizer()
self.down_proj_input_quantizer = TensorQuantizer()
self.down_proj_weight_quantizer = TensorQuantizer()

self._register_temp_attribute("_enable_weight_quantization", False)
self._register_temp_attribute("_down_proj_mul", False)
self._setup_for_weight_quantization()

@property
def functionals_to_replace(self):
# Use torch.ops.aten to bypass Python dispatch and avoid RecursionError
# (torch.matmul / __matmul__ can dispatch to each other)
_aten_bmm = torch.ops.aten.bmm
_aten_matmul = torch.ops.aten.matmul

def _quantized_bmm(batch1, batch2, *, out=None):
batch1 = self.down_proj_input_quantizer(batch1) if self._down_proj_mul else batch1
self._down_proj_mul = not self._down_proj_mul # toggle the flag
if out is not None:
return torch.ops.aten.bmm.out(batch1, batch2, out=out)
return _aten_bmm(batch1, batch2)

def _tensor_matmul(self_t, other):
self_t = self.down_proj_input_quantizer(self_t) if self._down_proj_mul else self_t
self._down_proj_mul = not self._down_proj_mul
return _aten_matmul(self_t, other)

return [
(torch, "bmm", _quantized_bmm),
(torch.Tensor, "__matmul__", _tensor_matmul),
]

@contextmanager
def quantize_weight(self):
"""Context in which MoE weight is quantized."""
self._enable_weight_quantization = True
try:
yield
finally:
for module in self.modules():
if isinstance(module, TensorQuantizer) and hasattr(module, "_cached_quant_val"):
delattr(module, "_cached_quant_val")
self._enable_weight_quantization = False

def forward(
self, hidden_states: torch.Tensor, router_indices=None, routing_weights=None
) -> torch.Tensor:
"""Forward method to add quantization."""
hidden_states = self.gate_up_proj_input_quantizer(hidden_states)
with self.quantize_weight():
return super().forward(hidden_states, router_indices, routing_weights)


if GptOssExperts not in QuantModuleRegistry:
QuantModuleRegistry.register({GptOssExperts: "hf.GptOssExperts"})(_QuantGptOssExperts)
92 changes: 1 addition & 91 deletions modelopt/torch/quantization/plugins/huggingface.py
Original file line number Diff line number Diff line change
Expand Up @@ -30,7 +30,6 @@
from torch import Tensor
from torch.nn.functional import linear
from transformers.integrations.finegrained_fp8 import FP8Linear
from transformers.models.gpt_oss.modeling_gpt_oss import GptOssExperts
from transformers.models.t5.modeling_t5 import T5Attention

from modelopt.torch.kernels.common.attention import IS_AVAILABLE as TRITON_FA_AVAILABLE
Expand Down Expand Up @@ -1206,96 +1205,6 @@ def unpack_weight(self):
QuantModuleRegistry.register({FP8Linear: "hf.FP8Linear"})(_QuantFP8Linear)


class _QuantGptOssExperts(_TransposedExpertsCalibMixin, _QuantFunctionalMixin):
"""Quantized wrapper for `transformers.GptOssExperts`.

Quantizes `gate_up_proj` and `down_proj` weights via dynamic attributes inside `quantize_weight()`.
Activations into `gate_up_proj` are quantized by `gate_up_proj_input_quantizer`. For `down_proj`
activation quantization, we intercept `torch.Tensor.__matmul__`/`torch.bmm` and quantize inputs
on every second call (since the first call computes `gate_up_proj` outputs and second call
computes `down_proj` outputs).
"""

@staticmethod
def _get_quantized_weight(quantizer, module, weight):
# MoE weight is accessed for each expert in one forward pass. so lets cache it
if module._enable_weight_quantization:
if hasattr(quantizer, "_cached_quant_val"):
return getattr(quantizer, "_cached_quant_val")
quantizer._cached_quant_val = _transposed_quantize(weight, quantizer)
return quantizer._cached_quant_val
return weight

def _setup_for_weight_quantization(self):
self._register_dynamic_attribute(
"gate_up_proj", partial(self._get_quantized_weight, self.gate_up_proj_weight_quantizer)
)
self._register_dynamic_attribute(
"down_proj", partial(self._get_quantized_weight, self.down_proj_weight_quantizer)
)

def _setup(self):
assert not hasattr(self, "kernel_layer_name"), (
"ModelOpt quantization does not support patched forward for kernel_hub"
)
self.gate_up_proj_input_quantizer = TensorQuantizer()
self.gate_up_proj_weight_quantizer = TensorQuantizer()
self.down_proj_input_quantizer = TensorQuantizer()
self.down_proj_weight_quantizer = TensorQuantizer()

self._register_temp_attribute("_enable_weight_quantization", False)
self._register_temp_attribute("_down_proj_mul", False)
self._setup_for_weight_quantization()

@property
def functionals_to_replace(self):
# Use torch.ops.aten to bypass Python dispatch and avoid RecursionError
# (torch.matmul / __matmul__ can dispatch to each other)
_aten_bmm = torch.ops.aten.bmm
_aten_matmul = torch.ops.aten.matmul

def _quantized_bmm(batch1, batch2, *, out=None):
batch1 = self.down_proj_input_quantizer(batch1) if self._down_proj_mul else batch1
self._down_proj_mul = not self._down_proj_mul # toggle the flag
if out is not None:
return torch.ops.aten.bmm.out(batch1, batch2, out=out)
return _aten_bmm(batch1, batch2)

def _tensor_matmul(self_t, other):
self_t = self.down_proj_input_quantizer(self_t) if self._down_proj_mul else self_t
self._down_proj_mul = not self._down_proj_mul
return _aten_matmul(self_t, other)

return [
(torch, "bmm", _quantized_bmm),
(torch.Tensor, "__matmul__", _tensor_matmul),
]

@contextmanager
def quantize_weight(self):
"""Context in which MoE weight is quantized."""
self._enable_weight_quantization = True
try:
yield
finally:
for module in self.modules():
if isinstance(module, TensorQuantizer) and hasattr(module, "_cached_quant_val"):
delattr(module, "_cached_quant_val")
self._enable_weight_quantization = False

def forward(
self, hidden_states: torch.Tensor, router_indices=None, routing_weights=None
) -> torch.Tensor:
"""Forward method to add quantization."""
hidden_states = self.gate_up_proj_input_quantizer(hidden_states)
with self.quantize_weight():
return super().forward(hidden_states, router_indices, routing_weights)


if GptOssExperts not in QuantModuleRegistry:
QuantModuleRegistry.register({GptOssExperts: "hf.GptOssExperts"})(_QuantGptOssExperts)


def _has_num_experts(obj):
# n_routed_experts: NemotronH-style MoE
return hasattr(obj, "num_experts") or hasattr(obj, "n_routed_experts")
Expand Down Expand Up @@ -1572,6 +1481,7 @@ def _is_param_grad_enabled_for_auto_quantize(pname, model):
# Nemotron-H's is the more specific one.
for _model_type in (
"falcon",
"gpt_oss",
"llama4",
"nemotron_h",
):
Expand Down
48 changes: 32 additions & 16 deletions tests/unit/torch/models/test_modeling_ptq_registration.py
Original file line number Diff line number Diff line change
Expand Up @@ -21,15 +21,18 @@
so an in-process test would only ever see the import order of whichever test ran first.
"""

import os
import subprocess
import sys
from concurrent.futures import ThreadPoolExecutor

import pytest

pytest.importorskip("transformers")

FIRST_IMPORTS = [
"modelopt.torch.models.falcon.modeling_ptq",
"modelopt.torch.models.gpt_oss.modeling_ptq",
"modelopt.torch.models.llama4.modeling_ptq",
"modelopt.torch.models.nemotron_h.modeling_ptq",
"modelopt.torch.quantization",
Expand Down Expand Up @@ -63,6 +66,13 @@
else:
assert Llama4TextExperts in QuantModuleRegistry, "Llama4TextExperts not registered"

try:
from transformers.models.gpt_oss.modeling_gpt_oss import GptOssExperts
except ImportError:
pass
else:
assert GptOssExperts in QuantModuleRegistry, "GptOssExperts not registered"

assert register_falcon_linears_on_the_fly in CUSTOM_MODEL_PLUGINS, "Falcon callback missing"

# The first matching decoder discoverer wins, so Nemotron-H's must precede the generic one.
Expand All @@ -73,23 +83,29 @@
"""


def _run_check(first_import):
# The child only asserts; tracing its heavy imports for coverage (inherited through
# COVERAGE_PROCESS_START) would just slow it down.
env = {k: v for k, v in os.environ.items() if k != "COVERAGE_PROCESS_START"}
result = subprocess.run(
[sys.executable, "-c", CHECK, first_import],
capture_output=True,
text=True,
env=env,
check=False,
)
return first_import, result


# Several fresh interpreters each import torch and transformers, which outlasts the default
# 60 s per-test cap on a small CI runner.
@pytest.mark.timeout(300)
def test_registration_is_independent_of_import_order():
# Each first import needs its own interpreter, but starting them one after another would
# add over a minute to the unit-test job, so launch them all and then collect the results.
procs = {
first_import: subprocess.Popen(
[sys.executable, "-c", CHECK, first_import],
stdout=subprocess.PIPE,
stderr=subprocess.PIPE,
text=True,
)
for first_import in FIRST_IMPORTS
}
failures = {}
for first_import, proc in procs.items():
_, stderr = proc.communicate(timeout=600)
if proc.returncode != 0:
failures[first_import] = stderr
# Each first import needs its own interpreter; run them in parallel, at most one per CPU,
# rather than one after another.
with ThreadPoolExecutor(max_workers=os.cpu_count() or 1) as pool:
results = list(pool.map(_run_check, FIRST_IMPORTS))
failures = {name: result.stderr for name, result in results if result.returncode != 0}
assert not failures, "\n\n".join(
f"importing {name} first:\n{err}" for name, err in failures.items()
)
Loading