Fix imports
This commit is contained in:
parent
8c668f492d
commit
c3055f48d7
10 changed files with 25 additions and 24 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -100,6 +100,7 @@ VLLM_SUPPORTED_VLM = [
|
|||
"gemma3",
|
||||
"mistral3",
|
||||
"qwen3_vl",
|
||||
"qwen3_vl_moe",
|
||||
]
|
||||
VLLM_NON_LORA_VLM = [
|
||||
"mllama",
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue