Building the Model Training Pipeline Backend - SFT, DPO, LoRA
This post documents the implementation of the training pipeline backend for the XGen platform, which lets users train custom LLMs. It supports three training methods: SFT (Supervised Fine-Tuning), DPO (Direct Preference Optimization), and LoRA (Low-Rank Adaptation).
Training Configuration Model
from pydantic import BaseModel, Field
from enum import Enum
class TrainingMethod(str, Enum):
SFT = "sft"
DPO = "dpo"
class TrainingConfig(BaseModel):
base_model: str = "meta-llama/Llama-3.1-8B"
method: TrainingMethod = TrainingMethod.SFT
use_lora: bool = True
lora_r: int = Field(default=16, ge=4, le=128)
lora_alpha: int = Field(default=32, ge=4, le=256)
lora_dropout: float = Field(default=0.05, ge=0.0, le=0.5)
learning_rate: float = Field(default=2e-5, ge=1e-7, le=1e-2)
num_epochs: int = Field(default=3, ge=1, le=100)
batch_size: int = Field(default=4, ge=1, le=64)
max_seq_length: int = Field(default=2048, ge=128, le=8192)
dataset_id: str
output_name: str
Dataset Processing
class DatasetProcessor:
@staticmethod
def prepare_sft_dataset(raw_data: list) -> list:
formatted = []
for item in raw_data:
formatted.append({
"instruction": item.get("instruction", ""),
"input": item.get("input", ""),
"output": item.get("output", ""),
})
return formatted
@staticmethod
def prepare_dpo_dataset(raw_data: list) -> list:
formatted = []
for item in raw_data:
formatted.append({
"prompt": item["prompt"],
"chosen": item["chosen"],
"rejected": item["rejected"],
})
return formatted
Training Job Management
Since training runs on VastAI GPU instances, remote job management was the central challenge.
class TrainingJob(Base):
__tablename__ = "training_jobs"
id = Column(String, primary_key=True)
user_id = Column(String, ForeignKey("users.id"))
config = Column(JSON, nullable=False)
status = Column(String(20), default="pending")
instance_id = Column(Integer, nullable=True)
progress = Column(Float, default=0.0)
metrics = Column(JSON, default=dict)
error_message = Column(String, nullable=True)
created_at = Column(DateTime, default=datetime.utcnow)
completed_at = Column(DateTime, nullable=True)
class TrainingService:
async def start_training(self, config: TrainingConfig, user_id: str) -> str:
job = TrainingJob(
id=str(uuid.uuid4()),
user_id=user_id,
config=config.model_dump(),
)
self.db.add(job)
await self.db.flush()
# Provision a GPU instance
instance = await self.gpu_service.provision_instance(
gpu_type="A100" if config.base_model.endswith("70B") else "RTX_4090",
image="pytorch/pytorch:2.3.0-cuda12.1-cudnn8-runtime",
)
job.instance_id = instance["id"]
job.status = "provisioning"
# Kick off training in the background
asyncio.create_task(self._run_training(job.id, config, instance))
return job.id
Progress Monitoring
Training progress is monitored in real time over WebSocket.
@router.websocket("/ws/training/{job_id}")
async def training_progress(websocket: WebSocket, job_id: str):
await websocket.accept()
try:
while True:
job = await training_service.get_job(job_id)
await websocket.send_json({
"status": job.status,
"progress": job.progress,
"metrics": job.metrics,
"epoch": job.metrics.get("current_epoch", 0),
"loss": job.metrics.get("loss", 0),
})
if job.status in ("completed", "failed"):
break
await asyncio.sleep(5)
finally:
await websocket.close()
API Endpoints
@router.post("/api/training/start")
async def start_training(config: TrainingConfig, user=Depends(auth_middleware)):
job_id = await training_service.start_training(config, user["user_id"])
return {"job_id": job_id}
@router.get("/api/training/{job_id}")
async def get_training_status(job_id: str):
job = await training_service.get_job(job_id)
return {"job": job}
The training pipeline was one of the most complex features in XGen. Automating the entire flow — GPU provisioning, training execution, monitoring, and model saving — required considerable effort.