Skip to content

Commit 9020480

Browse files
committed
Training Pipeline
1 parent 993e89f commit 9020480

21 files changed

Lines changed: 5135 additions & 6 deletions

deepfabric/__init__.py

Lines changed: 33 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -17,8 +17,32 @@
1717
from .generator import DataSetGenerator, DataSetGeneratorConfig
1818
from .graph import Graph, GraphConfig
1919
from .hf_hub import HFUploader
20+
from .pipeline import DeepFabricPipeline
2021
from .tree import Tree, TreeConfig
2122

23+
# Training module (optional dependencies)
24+
try:
25+
from .training import DeepFabricSFTTrainer, SFTTrainingConfig
26+
from .training.sft_config import LoRAConfig, QuantizationConfig
27+
28+
_has_training = True
29+
except ImportError:
30+
_has_training = False
31+
DeepFabricSFTTrainer = None
32+
SFTTrainingConfig = None
33+
LoRAConfig = None
34+
QuantizationConfig = None
35+
36+
# Transformers provider (optional dependencies)
37+
try:
38+
from .llm.transformers_provider import TransformersConfig, TransformersProvider
39+
40+
_has_transformers = True
41+
except ImportError:
42+
_has_transformers = False
43+
TransformersConfig = None
44+
TransformersProvider = None
45+
2246
__version__ = "0.1.0"
2347

2448
__all__ = [
@@ -31,7 +55,16 @@
3155
"Dataset",
3256
"DeepFabricConfig",
3357
"HFUploader",
58+
"DeepFabricPipeline",
3459
"cli",
60+
# Training (optional)
61+
"DeepFabricSFTTrainer",
62+
"SFTTrainingConfig",
63+
"LoRAConfig",
64+
"QuantizationConfig",
65+
# Transformers provider (optional)
66+
"TransformersConfig",
67+
"TransformersProvider",
3568
# Exceptions
3669
"DeepFabricError",
3770
"ConfigurationError",

deepfabric/format_command.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -52,7 +52,7 @@ def format_command(
5252
tui.info(f"Loading dataset from Hugging Face repo '{repo}' (split: {hf_split})...")
5353
try:
5454
# Bandit nosec, as no digest is set.
55-
hf_ds = load_dataset(str(repo), split=hf_split) # nosec
55+
hf_ds = load_dataset(str(repo), split=hf_split) # nosec
5656
except (DatasetNotFoundError, UnexpectedSplitsError) as e:
5757
msg = (
5858
"Failed to load dataset from Hugging Face repo "

deepfabric/llm/client.py

Lines changed: 45 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -156,6 +156,41 @@ def make_outlines_model(provider: str, model_name: str, **kwargs) -> Any:
156156
)
157157
return outlines.from_openai(client, model_name)
158158

159+
if provider == "transformers":
160+
# Use local HuggingFace Transformers model
161+
from .transformers_provider import TransformersProvider # noqa: PLC0415
162+
163+
# Create the provider and keep it alive
164+
transformers_provider = TransformersProvider(model_name, **kwargs)
165+
outlines_model = transformers_provider.get_outlines_model()
166+
167+
# Create a callable wrapper that uses the Transformers.generate method
168+
# This makes the transformers provider work like the API providers
169+
def transformers_callable(prompt: str, schema: type[BaseModel], **gen_kwargs):
170+
# Import generator module to create logits processor
171+
from outlines import generator # noqa: PLC0415
172+
173+
# Create JSON schema string from Pydantic model
174+
json_schema_str = schema.model_json_schema()
175+
# Get JSON string representation
176+
import json # noqa: PLC0415
177+
178+
json_schema_json = json.dumps(json_schema_str)
179+
180+
# Convert to logits processor for transformers
181+
logits_processor = generator.get_json_schema_logits_processor(
182+
None, # backend_name (None for auto-detect)
183+
outlines_model,
184+
json_schema_json,
185+
)
186+
# Generate using the model's generate method
187+
return outlines_model.generate(prompt, output_type=logits_processor, **gen_kwargs)
188+
189+
# Attach the provider to the callable to keep it alive
190+
transformers_callable._provider = transformers_provider # type: ignore[attr-defined]
191+
192+
return transformers_callable
193+
159194
_raise_unsupported_provider_error(provider)
160195

161196
except DataSetGeneratorError:
@@ -324,7 +359,12 @@ async def _generate_async_with_retry(self, prompt: str, schema: Any, **kwargs) -
324359
if self.provider == "gemini" and isinstance(schema, type) and issubclass(schema, BaseModel):
325360
generation_schema = _create_gemini_compatible_schema(schema)
326361

327-
json_output = await self.async_model(prompt, generation_schema, **kwargs)
362+
# Ensure async_model is available; fallback to synchronous generation in a thread if not.
363+
model = self.async_model
364+
if model is None:
365+
return await asyncio.to_thread(self.generate, prompt, schema, **kwargs)
366+
367+
json_output = await model(prompt, generation_schema, **kwargs)
328368
# Validate with original schema to ensure proper validation
329369
return schema.model_validate_json(json_output)
330370

@@ -334,6 +374,10 @@ def _convert_generation_params(self, **kwargs) -> dict:
334374
if self.provider == "gemini" and "max_tokens" in kwargs:
335375
kwargs["max_output_tokens"] = kwargs.pop("max_tokens")
336376

377+
# Convert max_tokens to max_new_tokens for Transformers
378+
if self.provider == "transformers" and "max_tokens" in kwargs:
379+
kwargs["max_new_tokens"] = kwargs.pop("max_tokens")
380+
337381
return kwargs
338382

339383
def __repr__(self) -> str:

deepfabric/llm/rate_limit_config.py

Lines changed: 52 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -193,11 +193,59 @@ class OllamaRateLimitConfig(RateLimitConfig):
193193
)
194194

195195

196+
class TransformersRateLimitConfig(RateLimitConfig):
197+
"""HuggingFace Transformers-specific rate limit configuration.
198+
199+
Local inference with Transformers doesn't face API rate limits, but
200+
may encounter hardware-related failures (CUDA OOM, generation errors).
201+
This config uses minimal retries focused on recoverable errors.
202+
"""
203+
204+
max_retries: int = Field(
205+
default=2,
206+
ge=0,
207+
le=5,
208+
description="Minimal retries for local model inference",
209+
)
210+
base_delay: float = Field(
211+
default=1.0,
212+
ge=0.1,
213+
le=10.0,
214+
description="Base delay for local inference retry",
215+
)
216+
max_delay: float = Field(
217+
default=10.0,
218+
ge=1.0,
219+
le=60.0,
220+
description="Max delay for local inference retry",
221+
)
222+
backoff_strategy: BackoffStrategy = Field(
223+
default=BackoffStrategy.LINEAR,
224+
description="Linear backoff for hardware issues",
225+
)
226+
jitter: bool = Field(
227+
default=False,
228+
description="No jitter needed for local inference",
229+
)
230+
respect_retry_after: bool = Field(
231+
default=False,
232+
description="No retry-after headers from local models",
233+
)
234+
retry_on_status_codes: set[int] = Field(
235+
default_factory=set,
236+
description="No HTTP status codes for local inference",
237+
)
238+
retry_on_exceptions: list[str] = Field(
239+
default_factory=lambda: ["cuda", "out of memory", "generation"],
240+
description="Exception keywords specific to local model inference",
241+
)
242+
243+
196244
def get_default_rate_limit_config(provider: str) -> RateLimitConfig:
197245
"""Get the default rate limit configuration for a provider.
198246
199247
Args:
200-
provider: Provider name (openai, anthropic, gemini, ollama)
248+
provider: Provider name (openai, anthropic, gemini, ollama, transformers)
201249
202250
Returns:
203251
Provider-specific rate limit configuration with sensible defaults
@@ -207,6 +255,7 @@ def get_default_rate_limit_config(provider: str) -> RateLimitConfig:
207255
"anthropic": AnthropicRateLimitConfig(),
208256
"gemini": GeminiRateLimitConfig(),
209257
"ollama": OllamaRateLimitConfig(),
258+
"transformers": TransformersRateLimitConfig(),
210259
}
211260
return configs.get(provider, RateLimitConfig())
212261

@@ -218,7 +267,7 @@ def create_rate_limit_config(
218267
"""Create a rate limit configuration from a dictionary.
219268
220269
Args:
221-
provider: Provider name (openai, anthropic, gemini, ollama)
270+
provider: Provider name (openai, anthropic, gemini, ollama, transformers)
222271
config_dict: Configuration parameters as dictionary
223272
224273
Returns:
@@ -235,6 +284,7 @@ def create_rate_limit_config(
235284
"anthropic": AnthropicRateLimitConfig,
236285
"gemini": GeminiRateLimitConfig,
237286
"ollama": OllamaRateLimitConfig,
287+
"transformers": TransformersRateLimitConfig,
238288
}
239289

240290
config_class = config_classes.get(provider, RateLimitConfig)

0 commit comments

Comments
 (0)