File size: 8,922 Bytes
5882d2f
 
 
 
 
16101d0
 
 
 
db3f481
 
2cc3878
db3f481
 
 
 
 
 
 
 
 
5882d2f
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
9a235f0
5882d2f
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
2f77cb5
 
9a235f0
5882d2f
 
2f77cb5
5882d2f
 
2f77cb5
 
5882d2f
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
cacdb69
 
 
5882d2f
cacdb69
 
 
 
 
 
 
5882d2f
cacdb69
95704a7
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
from typing import Dict, List, Optional, Any
import uuid
from pydantic import BaseModel, Field, HttpUrl, root_validator, validator, SecretStr
from enum import Enum
from datetime import datetime


class ImageRequest(BaseModel):
    prompt: str
    model: str = "black-forest-labs/flux-dev"
    enhance_prompt: bool = False
    provider: Optional[str] = "replicate"

    class Config:
        json_schema_extra = {
            "example": {
                "prompt": "Create a futuristic cityscape at sunset.",
                "model": "black-forest-labs/flux-dev",
                "enhance_prompt": True,
            }
        }


# Models
class TrainingStatus(str, Enum):
    STARTING = "starting"
    PROCESSING = "processing"
    SUCCEEDED = "succeeded"
    FAILED = "failed"
    CANCELED = "canceled"


class WebhookEvent(str, Enum):
    START = "start"
    OUTPUT = "output"
    LOGS = "logs"
    COMPLETED = "completed"


class AutocaptionOptions(BaseModel):
    enabled: bool = Field(False, description="Automatically caption images")
    prefix: Optional[str] = Field(None, description="Text to prepend to captions")
    suffix: Optional[str] = Field(None, description="Text to routerend to captions")


class WandbConfig(BaseModel):
    api_key: SecretStr = Field(..., description="Weights & Biases API key")
    project: str = Field("flux_train_replicate", description="W&B project name")
    entity: Optional[str] = Field(None, description="W&B entity name")
    run: Optional[str] = Field(None, description="W&B run name")
    sample_prompts: Optional[str] = Field(
        None, description="Newline-separated prompts for W&B samples"
    )
    sample_interval: int = Field(
        100, description="Step interval for sampling output images"
    )
    save_interval: int = Field(
        100, description="Step interval for saving intermediate weights"
    )


class TrainingInput(BaseModel):
    destination: str = Field(
        ...,
        description="Model destination in format owner/name",
        pattern=r"^[A-Za-z0-9_\-]+/[A-Za-z0-9_\-]+$",
    )
    model_owner: str = Field(..., description="Owner of the model to train")
    model_name: str = Field(..., description="Name of the model to train")
    version_id: str = Field(
        ..., description="Version ID of the model to train"
    )  # Optional field with default None

    # Required training parameters
    input_images: str = Field(
        ..., description="URL to zip file containing training data"
    )
    trigger_word: str = Field("TOK", description="Trigger word for training")

    # Optional training parameters
    steps: int = Field(1000, ge=3, le=6000, description="Number of training steps")
    lora_rank: int = Field(16, ge=1, le=128, description="LoRA rank value")

    hf_repo_id: Optional[str] = Field(None, description="HuggingFace repository ID")
    hf_token: Optional[SecretStr] = Field(None, description="HuggingFace token")

    # W&B integration
    wandb: Optional[WandbConfig] = Field(
        None, description="Weights & Biases configuration"
    )

    # Autocaption settings
    autocaption: AutocaptionOptions = Field(
        default_factory=AutocaptionOptions, description="Auto-captioning configuration"
    )

    # Advanced training parameters
    learning_rate: float = Field(0.0004, description="Learning rate")
    batch_size: int = Field(1, description="Batch size")
    resolution: str = Field(
        "512,768,1024", description="Image resolutions for training"
    )
    caption_dropout_rate: float = Field(
        0.05, ge=0, le=1, description="Caption dropout rate"
    )
    optimizer: str = Field("adamw8bit", description="Optimizer to use")
    cache_latents_to_disk: bool = Field(False, description="Cache latents to disk")
    layers_to_optimize_regex: Optional[str] = Field(
        None, description="Regex for layers to optimize"
    )
    gradient_checkpointing: bool = Field(
        False, description="Enable gradient checkpointing"
    )

    # Skip training option
    skip_training_and_use_pretrained_hf_lora_url: Optional[HttpUrl] = Field(
        None, description="URL to pretrained HF LoRA to use instead of training"
    )

    # Webhook configuration
    webhook: Optional[HttpUrl] = Field(
        None, description="Webhook URL for notifications"
    )
    webhook_events_filter: Optional[List[WebhookEvent]] = Field(
        None, description="Events that trigger webhooks"
    )

    def to_replicate_payload(self) -> Dict[str, Any]:
        """Convert to Replicate API payload format"""
        # Build inputs dictionary with all the training parameters
        input_dict = {
            "input_images": str(self.input_images),
            "steps": self.steps,
            "lora_rank": self.lora_rank,
            "learning_rate": self.learning_rate,
            "batch_size": self.batch_size,
            "resolution": self.resolution,
            "caption_dropout_rate": self.caption_dropout_rate,
            "optimizer": self.optimizer,
            "cache_latents_to_disk": self.cache_latents_to_disk,
            "gradient_checkpointing": self.gradient_checkpointing,
            "autocaption": self.autocaption.enabled,
        }

        # Add optional parameters if they exist
        if self.autocaption.prefix:
            input_dict["autocaption_prefix"] = self.autocaption.prefix
        if self.autocaption.suffix:
            input_dict["autocaption_suffix"] = self.autocaption.suffix
        if self.layers_to_optimize_regex:
            input_dict["layers_to_optimize_regex"] = self.layers_to_optimize_regex
        if self.hf_repo_id:
            input_dict["hf_repo_id"] = self.hf_repo_id
        if self.hf_token:
            input_dict["hf_token"] = self.hf_token.get_secret_value()
        if self.skip_training_and_use_pretrained_hf_lora_url:
            input_dict["skip_training_and_use_pretrained_hf_lora_url"] = str(
                self.skip_training_and_use_pretrained_hf_lora_url
            )

        # Handle W&B configuration
        if self.wandb:
            input_dict["wandb_api_key"] = self.wandb.api_key.get_secret_value()
            input_dict["wandb_project"] = self.wandb.project
            if self.wandb.entity:
                input_dict["wandb_entity"] = self.wandb.entity
            if self.wandb.run:
                input_dict["wandb_run"] = self.wandb.run
            if self.wandb.sample_prompts:
                input_dict["wandb_sample_prompts"] = self.wandb.sample_prompts
            input_dict["wandb_sample_interval"] = self.wandb.sample_interval
            input_dict["wandb_save_interval"] = self.wandb.save_interval

        # Build the final payload
        payload = {"destination": self.destination, "input": input_dict}

        # Add webhook configuration if present
        if self.webhook:
            payload["webhook"] = str(self.webhook)
            if self.webhook_events_filter:
                payload["webhook_events_filter"] = [
                    event.value for event in self.webhook_events_filter
                ]

        return payload


class TrainingInputExtended(TrainingInput):
    image_urls: Optional[List[HttpUrl]] = Field(
        None, description="List of image URLs to be zipped for training"
    )

    # Override the validator to allow either input_images or image_urls
    @root_validator(pre=True)
    def validate_input_source(cls, values):
        # Check if at least one of input_images or image_urls is provided
        if not values.get("input_images") and not values.get("image_urls"):
            raise ValueError("Either input_images or image_urls must be provided")
        return values


class TrainingResponse(BaseModel):
    id: str
    model: str
    version: str
    input: Dict[str, Any]
    logs: Optional[str] = None
    error: Optional[str] = None
    status: TrainingStatus
    created_at: datetime
    started_at: Optional[datetime] = None
    completed_at: Optional[datetime] = None
    output: Optional[Dict[str, Any]] = None
    metrics: Optional[Dict[str, float]] = None
    urls: Dict[str, str]


class PaginatedTrainings(BaseModel):
    next: Optional[str] = None
    previous: Optional[str] = None
    results: List[TrainingResponse]


class HardwareInfo(BaseModel):
    name: str
    sku: str


class ZipRequest(BaseModel):
    image_urls: List[str] = Field(
        ..., description="List of image URLs to include in the zip"
    )
    folder_id: Optional[str] = Field(
        default_factory=lambda: str(uuid.uuid4()), description="Unique folder ID"
    )
    caption: str = Field(
        "Describe this image", description="Caption instruction for image descriptions"
    )
    focus_entity: Optional[str] = Field(
        None, description="Focus instruction for image descriptions"
    )
    trigger_word: Optional[str] = Field(None, description="Trigger word for LoRA")

    image_urls_captions: Optional[List[str]] = None