Skip to content

Commit 96f5daa

Browse files
committed
Retry VLLM [release]
1 parent 4449d59 commit 96f5daa

3 files changed

Lines changed: 66 additions & 24 deletions

File tree

pyproject.toml

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,6 @@
11
[tool.poetry]
22
name = "DataDreamer"
3-
version = "0.44.0"
3+
version = "0.45.0"
44
description = "Prompt. Generate Synthetic Data. Train & Align Models."
55
license = "MIT"
66
authors= [
@@ -66,7 +66,7 @@ warn_unused_ignores = true
6666
mypy_path = "src/_stubs"
6767

6868
[[tool.mypy.overrides]]
69-
module = "click,wandb,wandb.*,click.testing,flaky,tensorflow,torch_xla,jax,datasets.features.features,datasets.iterable_dataset,datasets.fingerprint,datasets.builder,datasets.arrow_writer,datasets.splits,datasets.utils,datasets.utils.version,pyarrow.lib,huggingface_hub,huggingface_hub.utils._headers,huggingface_hub.errors,dill,dill.source,transformers,bitsandbytes,sqlitedict,optimum.bettertransformer,optimum.bettertransformer.models,optimum.utils,transformers.utils.quantization_config,sortedcontainers,peft,psutil,ring,ctransformers,petals,petals.client.inference_session,hivemind.p2p.p2p_daemon_bindings.utils,huggingface_hub.utils,tqdm,ctransformers.transformers,vllm,litellm,litellm.llms.palm,litellm.exceptions,sentence_transformers,faiss,huggingface_hub.utils._validators,evaluate,transformers.trainer_callback,transformers.training_args,trl,guidance,sentence_transformers.models.Transformer,trl.trainer.utils,transformers.trainer_utils,setfit,joblib,setfit.modeling,transformers.utils.notebook,mistralai.models,mistralai.models.chat_completion,accelerate.utils,accelerate.utils.constants,accelerate,transformers.trainer,sentence_transformers.util,Pyro5,Pyro5.server,Pyro5.api,Pyro5,datadreamer,huggingface_hub.repocard,transformers.trainer_pt_utils,traitlets.utils.warnings,orjson,Pyro5.errors,sympy,tqdm.auto"
69+
module = "click,wandb,wandb.*,click.testing,flaky,tensorflow,torch_xla,jax,datasets.features.features,datasets.iterable_dataset,datasets.fingerprint,datasets.builder,datasets.arrow_writer,datasets.splits,datasets.utils,datasets.utils.version,pyarrow.lib,huggingface_hub,huggingface_hub.utils._headers,huggingface_hub.errors,dill,dill.source,transformers,bitsandbytes,sqlitedict,optimum.bettertransformer,optimum.bettertransformer.models,optimum.utils,transformers.utils.quantization_config,sortedcontainers,peft,psutil,ring,ctransformers,petals,petals.client.inference_session,hivemind.p2p.p2p_daemon_bindings.utils,huggingface_hub.utils,tqdm,ctransformers.transformers,vllm,litellm,litellm.llms.palm,litellm.exceptions,sentence_transformers,faiss,huggingface_hub.utils._validators,evaluate,transformers.trainer_callback,transformers.training_args,trl,guidance,sentence_transformers.models.Transformer,trl.trainer.utils,transformers.trainer_utils,setfit,joblib,setfit.modeling,transformers.utils.notebook,mistralai.models,mistralai.models.chat_completion,accelerate.utils,accelerate.utils.constants,accelerate,transformers.trainer,sentence_transformers.util,Pyro5,Pyro5.server,Pyro5.api,Pyro5,datadreamer,huggingface_hub.repocard,transformers.trainer_pt_utils,traitlets.utils.warnings,orjson,Pyro5.errors,sympy,tqdm.auto,requests.exceptions"
7070
ignore_missing_imports = true
7171

7272
[tool.pyright]

src/llms/vllm.py

Lines changed: 58 additions & 15 deletions
Original file line numberDiff line numberDiff line change
@@ -6,6 +6,14 @@
66
from typing import Any, Callable, Generator, Iterable
77

88
import torch
9+
from tenacity import (
10+
after_log,
11+
before_sleep_log,
12+
retry,
13+
retry_if_exception_type,
14+
wait_exponential,
15+
)
16+
917
from datasets.fingerprint import Hasher
1018

1119
from .. import DataDreamer
@@ -72,6 +80,34 @@ def __init__(
7280
self.quantization = "awq"
7381
self.swap_space = swap_space
7482

83+
@cached_property
84+
def retry_wrapper(self):
85+
# Create a retry wrapper function
86+
tenacity_logger = self.get_logger(key="retry", verbose=True, log_level=None)
87+
88+
from requests.exceptions import ConnectionError, Timeout
89+
90+
@retry(
91+
retry=retry_if_exception_type(Timeout),
92+
wait=wait_exponential(multiplier=1, min=10, max=60),
93+
before_sleep=before_sleep_log(tenacity_logger, logging.INFO),
94+
after=after_log(tenacity_logger, logging.INFO),
95+
reraise=True,
96+
)
97+
@retry(
98+
retry=retry_if_exception_type(ConnectionError),
99+
wait=wait_exponential(multiplier=1, min=10, max=60),
100+
before_sleep=before_sleep_log(tenacity_logger, logging.INFO),
101+
after=after_log(tenacity_logger, logging.INFO),
102+
reraise=True,
103+
)
104+
def _retry_wrapper(func, *args, **kwargs):
105+
return func(*args, **kwargs)
106+
107+
_retry_wrapper.__wrapped__.__module__ = None # type: ignore[attr-defined]
108+
_retry_wrapper.__wrapped__.__qualname__ = f"{self.__class__.__name__}.run" # type: ignore[attr-defined]
109+
return _retry_wrapper
110+
75111
@cached_property
76112
def model(self) -> Any:
77113
env = os.environ.copy()
@@ -119,17 +155,21 @@ def _monkey_patch_init_logger(*args, **kwargs):
119155
if datadreamer_logger.level > logging.DEBUG
120156
else nullcontext()
121157
):
122-
self_resource.model = LLM(
123-
model=self.model_name,
124-
trust_remote_code=self.trust_remote_code,
125-
dtype=str(self.dtype).replace("torch.", "")
126-
if self.dtype is not None
127-
else "auto",
128-
quantization=self.quantization,
129-
revision=self.revision,
130-
swap_space=self.swap_space,
131-
tensor_parallel_size=tensor_parallel_size,
132-
**kwargs,
158+
self_resource.model = self.retry_wrapper(
159+
func=(
160+
lambda: LLM(
161+
model=self.model_name,
162+
trust_remote_code=self.trust_remote_code,
163+
dtype=str(self.dtype).replace("torch.", "")
164+
if self.dtype is not None
165+
else "auto",
166+
quantization=self.quantization,
167+
revision=self.revision,
168+
swap_space=self.swap_space,
169+
tensor_parallel_size=tensor_parallel_size,
170+
**kwargs,
171+
)
172+
)
133173
)
134174

135175
# Finished loading
@@ -144,7 +184,9 @@ def _monkey_patch_init_logger(*args, **kwargs):
144184

145185
@dill_serializer
146186
def get_generated_texts_batch(self_resource, *args, **kwargs):
147-
outputs = self_resource.model.generate(*args, **kwargs)
187+
outputs = self.retry_wrapper(
188+
self_resource.model.generate, *args, **kwargs
189+
)
148190
generated_texts_batch = [
149191
[o.text for o in batch.outputs] for batch in outputs
150192
]
@@ -183,9 +225,9 @@ def _run_batch( # noqa: C901
183225
**kwargs,
184226
) -> list[str] | list[list[str]]:
185227
prompts = inputs
186-
assert (
187-
logit_bias is None
188-
), f"`logit_bias` is not supported for {type(self).__name__}"
228+
assert logit_bias is None, (
229+
f"`logit_bias` is not supported for {type(self).__name__}"
230+
)
189231
assert seed is None, f"`seed` is not supported for {type(self).__name__}"
190232

191233
SamplingParams = import_module("vllm").SamplingParams
@@ -323,6 +365,7 @@ def __getstate__(self): # pragma: no cover
323365
state = super().__getstate__()
324366

325367
# Remove cached model or tokenizer before serializing
368+
state.pop("retry_wrapper", None)
326369
state.pop("model", None)
327370
state.pop("tokenizer", None)
328371

src/trainers/_train_hf_base.py

Lines changed: 6 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -7,7 +7,6 @@
77
from typing import TYPE_CHECKING, Any, Type, cast
88

99
import torch
10-
1110
from datasets.fingerprint import Hasher
1211

1312
from .. import DataDreamer
@@ -490,9 +489,9 @@ def export_to_disk(self, path: str, adapter_only: bool = False) -> PreTrainedMod
490489
from .train_hf_finetune import TrainHFFineTune
491490
from .train_setfit_classifier import TrainSetFitClassifier
492491

493-
assert not adapter_only or self.peft_config, (
494-
"`adapter_only` can only be used if a `peft_config` was provided."
495-
)
492+
assert (
493+
not adapter_only or self.peft_config
494+
), "`adapter_only` can only be used if a `peft_config` was provided."
496495

497496
# Clear the directory
498497
clear_dir(path)
@@ -625,9 +624,9 @@ def publish_to_hf_hub( # noqa: C901
625624
from .train_hf_finetune import TrainHFFineTune
626625
from .train_setfit_classifier import TrainSetFitClassifier
627626

628-
assert not adapter_only or self.peft_config, (
629-
"`adapter_only` can only be used if a `peft_config` was provided."
630-
)
627+
assert (
628+
not adapter_only or self.peft_config
629+
), "`adapter_only` can only be used if a `peft_config` was provided."
631630

632631
# Login
633632
api = hf_hub_login(token=token)

0 commit comments

Comments
 (0)