Fix imports

This commit is contained in:
Datta Nimmaturi 2026-01-06 06:46:38 +00:00
commit c3055f48d7
10 changed files with 25 additions and 24 deletions

View file

@ -158,7 +158,7 @@ def get_or_autotune_moe_kernels(
if cached_data is not None:
# Reconstruct config objects from cached data
try:
from grouped_gemm.kernels.tuning import (
from .grouped_gemm.kernels.tuning import (
KernelConfigForward,
KernelConfigBackward_dX,
KernelConfigBackward_dW,
@ -266,17 +266,17 @@ def _run_moe_autotuning(
# Autotune forward kernel - use the interface function with autotune=True
# This properly invokes the kernel and lets triton handle the autotuning
from grouped_gemm.interface import (
from .grouped_gemm.interface import (
grouped_gemm_forward,
grouped_gemm_dX,
grouped_gemm_dW,
)
from grouped_gemm.kernels.forward import _autotuned_grouped_gemm_forward_kernel
from grouped_gemm.kernels.backward import (
from .grouped_gemm.kernels.forward import _autotuned_grouped_gemm_forward_kernel
from .grouped_gemm.kernels.backward import (
_autotuned_grouped_gemm_dX_kernel,
_autotuned_grouped_gemm_dW_kernel,
)
from grouped_gemm.kernels.tuning import (
from .grouped_gemm.kernels.tuning import (
KernelConfigForward,
KernelConfigBackward_dX,
KernelConfigBackward_dW,
@ -368,7 +368,7 @@ def _run_moe_autotuning(
def _get_default_configs() -> Tuple[Any, Any, Any]:
"""Get default kernel configurations as fallback."""
from grouped_gemm.kernels.tuning import (
from .grouped_gemm.kernels.tuning import (
KernelConfigForward,
KernelConfigBackward_dX,
KernelConfigBackward_dW,

View file

@ -8,17 +8,17 @@ from dataclasses import asdict
import torch
import triton
from grouped_gemm.kernels.backward import (
from .kernels.backward import (
_autotuned_grouped_gemm_dW_kernel,
_autotuned_grouped_gemm_dX_kernel,
_grouped_gemm_dW_kernel,
_grouped_gemm_dX_kernel,
)
from grouped_gemm.kernels.forward import (
from .kernels.forward import (
_autotuned_grouped_gemm_forward_kernel,
_grouped_gemm_forward_kernel,
)
from grouped_gemm.kernels.tuning import (
from .kernels.tuning import (
KernelConfigBackward_dW,
KernelConfigBackward_dX,
KernelConfigForward,

View file

@ -323,8 +323,8 @@ def exceeds_smem_capacity(
def common_prune_criteria(config: triton.Config, kwargs: dict, dtype):
from grouped_gemm.interface import supports_tma
from grouped_gemm.kernels.tuning import get_device_properties
from ..interface import supports_tma
from .tuning import get_device_properties
smem_size = get_device_properties().SIZE_SMEM
@ -355,7 +355,7 @@ def common_prune_criteria(config: triton.Config, kwargs: dict, dtype):
def maybe_disable_tma(config: triton.Config):
from grouped_gemm.interface import supports_tma
from ..interface import supports_tma
tma_keys = [k for k in config.kwargs.keys() if k.startswith("USE_TMA_")]
if not supports_tma():

View file

@ -5,7 +5,7 @@ import torch
import triton
import triton.language as tl
from grouped_gemm.kernels.autotuning import (
from .autotuning import (
get_dW_kernel_configs,
get_dX_kernel_configs,
prune_dX_configs,

View file

@ -5,7 +5,7 @@ import torch
import triton
import triton.language as tl
from grouped_gemm.kernels.autotuning import (
from .autotuning import (
get_forward_configs,
prune_kernel_configs_fwd,
)

View file

@ -15,7 +15,7 @@ import torch
import triton
from triton.runtime.errors import OutOfResources
from grouped_gemm.kernels.autotuning import (
from .autotuning import (
BOOLS,
DEFAULT_K_BLOCK_SIZES,
DEFAULT_M_BLOCK_SIZES,

View file

@ -9,13 +9,13 @@ import torch.nn.functional as F
from transformers.models.llama4 import Llama4TextConfig
from transformers.models.llama4.modeling_llama4 import Llama4TextMoe
from grouped_gemm.interface import grouped_gemm
from grouped_gemm.kernels.tuning import (
from ...interface import grouped_gemm
from ...kernels.tuning import (
KernelConfigBackward_dW,
KernelConfigBackward_dX,
KernelConfigForward,
)
from grouped_gemm.reference.moe_ops import (
from ..moe_ops import (
get_routing_indices,
permute,
torch_grouped_gemm,

View file

@ -12,13 +12,13 @@ from transformers.models.qwen3_moe.modeling_qwen3_moe import (
Qwen3MoeSparseMoeBlock,
)
from grouped_gemm.interface import grouped_gemm
from grouped_gemm.kernels.tuning import (
from ...interface import grouped_gemm
from ...kernels.tuning import (
KernelConfigBackward_dW,
KernelConfigBackward_dX,
KernelConfigForward,
)
from grouped_gemm.reference.moe_ops import (
from ..moe_ops import (
get_routing_indices,
permute,
torch_grouped_gemm,

View file

@ -5,13 +5,13 @@ import torch
from transformers.models.qwen3_moe.configuration_qwen3_moe import Qwen3MoeConfig
from transformers.models.qwen3_moe.modeling_qwen3_moe import Qwen3MoeSparseMoeBlock
from grouped_gemm.interface import grouped_gemm
from grouped_gemm.kernels.tuning import (
from ..interface import grouped_gemm
from ..kernels.tuning import (
KernelConfigBackward_dW,
KernelConfigBackward_dX,
KernelConfigForward,
)
from grouped_gemm.reference.moe_ops import (
from .moe_ops import (
Qwen3MoeGroupedGEMMBlock,
permute,
unpermute,

View file

@ -100,6 +100,7 @@ VLLM_SUPPORTED_VLM = [
"gemma3",
"mistral3",
"qwen3_vl",
"qwen3_vl_moe",
]
VLLM_NON_LORA_VLM = [
"mllama",