Documents
Home>Documents>AI>LLM>Train & Tune

Building a Training Pipeline Backend: SFT, DPO, and LoRA

8 min readJun 25, 2025Feb 22, 2026

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.

Tags
TrainingSFTDPOLoRAFineTuning