diff --git a/unsloth/models/rl.py b/unsloth/models/rl.py index 0605467203..fa95a809b2 100644 --- a/unsloth/models/rl.py +++ b/unsloth/models/rl.py @@ -43,31 +43,10 @@ torch_compile_options = { "triton.cudagraphs": False, } -# vLLM compatibility shim (TRL expects GuidedDecodingParams even if vLLM doesn't provide it) -try: - import vllm.sampling_params as _unsloth_vllm_sp - - if not hasattr(_unsloth_vllm_sp, "GuidedDecodingParams"): - - class GuidedDecodingParams: - def __init__(self, **kwargs): - self.kwargs = kwargs - - _unsloth_vllm_sp.GuidedDecodingParams = GuidedDecodingParams -except Exception: - pass - -from trl import __version__ as trl_version_raw -from importlib.metadata import version as importlib_version +from trl import __version__ as trl_version from unsloth_zoo.utils import Version -try: - trl_version = Version(trl_version_raw) -except Exception: - try: - trl_version = Version(importlib_version("trl")) - except Exception: - trl_version = Version("0.0.0") +trl_version = Version(trl_version) def vLLMSamplingParams(**kwargs): @@ -241,7 +220,7 @@ RLTrainer_replacement = ''' import os from typing import * from dataclasses import dataclass, field -from unsloth_zoo.utils import Version +from packaging.version import Version import torch import numpy as np from contextlib import nullcontext