SQLModel: Modern Python ORM for Type-Safe AI Applications
In AI/ML application development, database management is a crucial component. From storing experiment results, managing model registries, to dataset versioning tracking, everything requires efficient and type-safe database interactions. SQLModel, created by Sebastian Ramirez (the creator of FastAPI), combines the power of SQLAlchemy and Pydantic to deliver the best ORM experience in Python.
In this tutorial, we will explore SQLModel in depth and build a production-ready ML experiment tracking and model registry system.
What Is SQLModel?
SQLModel is a Python library that combines SQLAlchemy (the most popular ORM in Python) with Pydantic (a data validation library). The result is an ORM that is:
- Type-safe: Full type hints and IDE autocompletion
- Auto-validated: Data is validated before entering the database
- FastAPI compatible: Seamless integration from the same creator
- Simple: One class for both database model AND API schema
- Powerful: Full access to SQLAlchemy features when needed
Installation and Setup
Installing SQLModel
pip install sqlmodel
For specific databases:
# PostgreSQL
pip install sqlmodel psycopg2-binary
MySQL
pip install sqlmodel pymysql
Async support
pip install sqlmodel aiosqlite asyncpg
Verify
python -c "import sqlmodel; print(sqlmodel.version)"
Database Engine Setup
from sqlmodel import createengine, SQLModel
SQLite (development)
sqlite
url = "sqlite:///./mltracking.db"
engine = create
engine(sqliteurl, echo=True)
PostgreSQL (production)
postgres
url = "postgresql://user:password@localhost:5432/mltracking"
engine = create
engine(postgresurl, echo=False, poolsize=20, maxoverflow=10)
Create all tables
SQLModel.metadata.create
all(engine)
Model Definition
Basic Models
from sqlmodel import Field, SQLModel, Session, select
from typing import Optional
from datetime import datetime
import uuid
class Experiment(SQLModel, table=True):
"""Model for ML Experiment"""
id: Optional[int] = Field(default=None, primarykey=True)
name: str = Field(index=True, minlength=1, maxlength=255)
description: Optional[str] = Field(default=None, maxlength=1000)
datasetname: str = Field(index=True)
algorithm: str
hyperparameters: Optional[str] = Field(default=None) # JSON string
accuracy: Optional[float] = Field(default=None, ge=0.0, le=1.0)
loss: Optional[float] = Field(default=None, ge=0.0)
status: str = Field(default="created", index=True)
createdat: datetime = Field(defaultfactory=datetime.utcnow)
updatedat: Optional[datetime] = Field(default=None)
createdby: str = Field(default="system")
class MLModel(SQLModel, table=True):
"""Model for Model Registry"""
tablename = "mlmodels"
id: Optional[int] = Field(default=None, primarykey=True)
name: str = Field(index=True, unique=True)
version: str = Field(default="1.0.0")
framework: str # pytorch, tensorflow, sklearn, etc.
modelpath: str
filesizemb: Optional[float] = Field(default=None)
inputschema: Optional[str] = Field(default=None)
outputschema: Optional[str] = Field(default=None)
isactive: bool = Field(default=True, index=True)
createdat: datetime = Field(defaultfactory=datetime.utcnow)
class Dataset(SQLModel, table=True):
"""Model for Dataset Registry"""
id: Optional[int] = Field(default=None, primarykey=True)
name: str = Field(index=True)
version: str = Field(default="1.0")
filepath: str
rowcount: Optional[int] = Field(default=None)
columncount: Optional[int] = Field(default=None)
fileformat: str = Field(default="csv")
description: Optional[str] = Field(default=None)
createdat: datetime = Field(defaultfactory=datetime.utcnow)
Models with Read and Create Schemas
SQLModel allows creating separate schemas for read and create operations:
from sqlmodel import SQLModel, Field
from typing import Optional
from datetime import datetime
Base model (shared fields)
class ExperimentBase(SQLModel):
name: str = Field(minlength=1, maxlength=255)
description: Optional[str] = None
datasetname: str
algorithm: str
hyperparameters: Optional[str] = None
Create schema (for API input)
class ExperimentCreate(ExperimentBase):
pass
Update schema
class ExperimentUpdate(SQLModel):
name: Optional[str] = None
description: Optional[str] = None
accuracy: Optional[float] = None
loss: Optional[float] = None
status: Optional[str] = None
Database model (table)
class Experiment(ExperimentBase, table=True):
id: Optional[int] = Field(default=None, primarykey=True)
accuracy: Optional[float] = Field(default=None, ge=0.0, le=1.0)
loss: Optional[float] = Field(default=None, ge=0.0)
status: str = Field(default="created", index=True)
createdat: datetime = Field(defaultfactory=datetime.utcnow)
updatedat: Optional[datetime] = None
Read schema (for API response)
class ExperimentRead(ExperimentBase):
id: int
accuracy: Optional[float]
loss: Optional[float]
status: str
createdat: datetime
CRUD Operations
Create (Insert Data)
from sqlmodel import Session
def createexperiment(engine, experimentdata: ExperimentCreate) -> Experiment:
"""Create a new experiment"""
experiment = Experiment.modelvalidate(experimentdata)
with Session(engine) as session:
session.add(experiment)
session.commit()
session.refresh(experiment)
return experiment
Usage example
newexp = ExperimentCreate(
name="ResNet50 Transfer Learning",
description="Fine-tuning ResNet50 on custom dataset",
datasetname="chest-xray-v2",
algorithm="ResNet50",
hyperparameters='{"lr": 0.001, "epochs": 50, "batchsize": 32}',
)
experiment = createexperiment(engine, newexp)
print(f"Experiment created: ID={experiment.id}")
Batch Insert
def createexperimentsbatch(engine, experiments: list[ExperimentCreate]) -> list[Experiment]:
"""Insert multiple experiments at once"""
db
experiments = [Experiment.modelvalidate(exp) for exp in experiments]
with Session(engine) as session:
session.add
all(dbexperiments)
session.commit()
for exp in db
experiments:
session.refresh(exp)
return dbexperiments
Read (Query Data)
from sqlmodel import Session, select
def getexperimentbyid(engine, experimentid: int) -> Optional[Experiment]:
"""Get experiment by ID"""
with Session(engine) as session:
return session.get(Experiment, experimentid)
def getallexperiments(engine, skip: int = 0, limit: int = 100) -> list[Experiment]:
"""Get all experiments with pagination"""
with Session(engine) as session:
statement = select(Experiment).offset(skip).limit(limit)
return session.exec(statement).all()
def getexperimentsbystatus(engine, status: str) -> list[Experiment]:
"""Filter experiments by status"""
with Session(engine) as session:
statement = select(Experiment).where(Experiment.status == status)
return session.exec(statement).all()
def getbestexperiments(engine, topn: int = 10) -> list[Experiment]:
"""Get experiments with highest accuracy"""
with Session(engine) as session:
statement = (
select(Experiment)
.where(Experiment.accuracy.isnot(None))
.orderby(Experiment.accuracy.desc())
.limit(topn)
)
return session.exec(statement).all()
def searchexperiments(engine, query: str) -> list[Experiment]:
"""Search experiments by name or algorithm"""
with Session(engine) as session:
statement = select(Experiment).where(
(Experiment.name.contains(query)) |
(Experiment.algorithm.contains(query))
)
return session.exec(statement).all()
Update
def updateexperiment(
engine, experiment
id: int, updatedata: ExperimentUpdate
) -> Optional[Experiment]:
"""Update an experiment"""
with Session(engine) as session:
experiment = session.get(Experiment, experiment
id)
if not experiment:
return None
updatedict = updatedata.modeldump(excludeunset=True)
for key, value in updatedict.items():
setattr(experiment, key, value)
experiment.updatedat = datetime.utcnow()
session.add(experiment)
session.commit()
session.refresh(experiment)
return experiment
Example: update training results
update = ExperimentUpdate(accuracy=0.95, loss=0.12, status="completed")
updated = updateexperiment(engine, experimentid=1, updatedata=update)
Delete
def deleteexperiment(engine, experimentid: int) -> bool:
"""Delete an experiment"""
with Session(engine) as session:
experiment = session.get(Experiment, experiment
id)
if not experiment:
return False
session.delete(experiment)
session.commit()
return True
def deleteexperimentsbystatus(engine, status: str) -> int:
"""Delete all experiments with a specific status"""
with Session(engine) as session:
statement = select(Experiment).where(Experiment.status == status)
experiments = session.exec(statement).all()
count = len(experiments)
for exp in experiments:
session.delete(exp)
session.commit()
return count
Relationships
SQLModel supports table relationships just like SQLAlchemy:
from sqlmodel import SQLModel, Field, Relationship
from typing import Optional
class Team(SQLModel, table=True):
id: Optional[int] = Field(default=None, primarykey=True)
name: str = Field(index=True, unique=True)
description: Optional[str] = None
# Relationship: one team has many experiments
experiments: list["Experiment"] = Relationship(backpopulates="team")
members: list["TeamMember"] = Relationship(backpopulates="team")
class TeamMember(SQLModel, table=True):
tablename = "teammembers"
id: Optional[int] = Field(default=None, primarykey=True)
name: str
email: str = Field(unique=True)
role: str = Field(default="member")
teamid: Optional[int] = Field(default=None, foreignkey="team.id")
team: Optional[Team] = Relationship(backpopulates="members")
class Experiment(SQLModel, table=True):
id: Optional[int] = Field(default=None, primarykey=True)
name: str = Field(index=True)
algorithm: str
accuracy: Optional[float] = None
status: str = Field(default="created")
createdat: datetime = Field(defaultfactory=datetime.utcnow)
# Foreign key to team
teamid: Optional[int] = Field(default=None, foreignkey="team.id")
team: Optional[Team] = Relationship(backpopulates="experiments")
# Relationship to metrics
metrics: list["ExperimentMetric"] = Relationship(backpopulates="experiment")
class ExperimentMetric(SQLModel, table=True):
tablename = "experimentmetrics"
id: Optional[int] = Field(default=None, primarykey=True)
epoch: int
trainloss: float
valloss: float
trainaccuracy: float
valaccuracy: float
experimentid: int = Field(foreignkey="experiment.id")
experiment: Optional[Experiment] = Relationship(backpopulates="metrics")
Querying with Relationships
def getteamwithexperiments(engine, teamid: int) -> Optional[Team]:
"""Get team with all its experiments"""
with Session(engine) as session:
statement = select(Team).where(Team.id == team
id)
team = session.exec(statement).first()
if team:
# Access experiments (lazy loaded)
= team.experiments
return team
def getexperimentwithmetrics(engine, experimentid: int):
"""Get experiment with all epoch metrics"""
with Session(engine) as session:
statement = select(Experiment).where(Experiment.id == experimentid)
experiment = session.exec(statement).first()
if experiment:
= experiment.metrics
return experiment
Integration with FastAPI
SQLModel is designed to work perfectly with FastAPI:
from fastapi import FastAPI, HTTPException, Depends, Query
from sqlmodel import Session, select, createengine, SQLModel
from typing import Optional
DATABASEURL = "postgresql://user:password@localhost:5432/mltracking"
engine = createengine(DATABASEURL, poolsize=20)
app = FastAPI(title="ML Experiment Tracker API")
def getsession():
"""Dependency for database session"""
with Session(engine) as session:
yield session
@app.onevent("startup")
def onstartup():
SQLModel.metadata.createall(engine)
@app.post("/experiments/", responsemodel=ExperimentRead)
def createexperiment(
experiment: ExperimentCreate,
session: Session = Depends(getsession),
):
dbexperiment = Experiment.modelvalidate(experiment)
session.add(dbexperiment)
session.commit()
session.refresh(dbexperiment)
return dbexperiment
@app.get("/experiments/", responsemodel=list[ExperimentRead])
def listexperiments(
skip: int = Query(default=0, ge=0),
limit: int = Query(default=20, le=100),
status: Optional[str] = None,
algorithm: Optional[str] = None,
session: Session = Depends(getsession),
):
statement = select(Experiment)
if status:
statement = statement.where(Experiment.status == status)
if algorithm:
statement = statement.where(Experiment.algorithm == algorithm)
statement = statement.offset(skip).limit(limit)
return session.exec(statement).all()
@app.get("/experiments/{experimentid}", responsemodel=ExperimentRead)
def getexperiment(experimentid: int, session: Session = Depends(getsession)):
experiment = session.get(Experiment, experimentid)
if not experiment:
raise HTTPException(statuscode=404, detail="Experiment not found")
return experiment
@app.patch("/experiments/{experimentid}", responsemodel=ExperimentRead)
def updateexperiment(
experimentid: int,
updatedata: ExperimentUpdate,
session: Session = Depends(getsession),
):
experiment = session.get(Experiment, experimentid)
if not experiment:
raise HTTPException(statuscode=404, detail="Experiment not found")
for key, value in updatedata.modeldump(excludeunset=True).items():
setattr(experiment, key, value)
experiment.updatedat = datetime.utcnow()
session.add(experiment)
session.commit()
session.refresh(experiment)
return experiment
@app.delete("/experiments/{experimentid}")
def deleteexperiment(experimentid: int, session: Session = Depends(getsession)):
experiment = session.get(Experiment, experimentid)
if not experiment:
raise HTTPException(statuscode=404, detail="Experiment not found")
session.delete(experiment)
session.commit()
return {"message": "Experiment deleted", "id": experimentid}
@app.get("/experiments/best/", responsemodel=list[ExperimentRead])
def getbestexperiments(
topn: int = Query(default=10, le=50),
session: Session = Depends(getsession),
):
statement = (
select(Experiment)
.where(Experiment.accuracy.isnot(None))
.orderby(Experiment.accuracy.desc())
.limit(topn)
)
return session.exec(statement).all()
Async Support
SQLModel supports asynchronous operations for high performance:
from sqlmodel import SQLModel
from sqlmodel.ext.asyncio.session import AsyncSession
from sqlalchemy.ext.asyncio import createasyncengine, AsyncEngine
from sqlalchemy.orm import sessionmaker
Async engine
asyncengine = createasyncengine(
"postgresql+asyncpg://user:password@localhost:5432/mltracking",
echo=False,
poolsize=20,
)
asyncsession = sessionmaker(asyncengine, class=AsyncSession, expireoncommit=False)
async def getasyncsession():
async with asyncsession() as session:
yield session
Async CRUD operations
from fastapi import FastAPI, Depends
from sqlmodel import select
app = FastAPI()
@app.post("/experiments/", responsemodel=ExperimentRead)
async def createexperimentasync(
experiment: ExperimentCreate,
session: AsyncSession = Depends(getasyncsession),
):
dbexperiment = Experiment.modelvalidate(experiment)
session.add(dbexperiment)
await session.commit()
await session.refresh(dbexperiment)
return dbexperiment
@app.get("/experiments/", responsemodel=list[ExperimentRead])
async def listexperimentsasync(
session: AsyncSession = Depends(getasyncsession),
):
result = await session.exec(select(Experiment))
return result.all()
Migration with Alembic
Alembic is the standard tool for database migrations in the SQLAlchemy/SQLModel ecosystem:
# Install
pip install alembic
Initialize
alembic init alembic
Configure alembic/env.py:
from sqlmodel import SQLModel
from app.models import Experiment, MLModel, Dataset, Team # Import all models
targetmetadata = SQLModel.metadata
def runmigrationsonline():
connectable = enginefromconfig(
config.getsection(config.configinisection),
prefix="sqlalchemy.",
)
with connectable.connect() as connection:
context.configure(
connection=connection,
targetmetadata=targetmetadata,
)
with context.begintransaction():
context.runmigrations()
Creating and running migrations:
# Generate automatic migration
alembic revision --autogenerate -m "add experiment metrics table"
Run migration
alembic upgrade head
Rollback
alembic downgrade -1
View history
alembic history
Advanced Query Building
from sqlmodel import Session, select, func, col, or, and
def advancedqueries(engine):
with Session(engine) as session:
# Aggregation: average accuracy per algorithm
statement = (
select(
Experiment.algorithm,
func.count(Experiment.id).label("total"),
func.avg(Experiment.accuracy).label("avgaccuracy"),
func.max(Experiment.accuracy).label("bestaccuracy"),
)
.where(Experiment.status == "completed")
.groupby(Experiment.algorithm)
.orderby(func.avg(Experiment.accuracy).desc())
)
results = session.exec(statement).all()
# Subquery: experiments above average
avgsubquery = select(func.avg(Experiment.accuracy)).scalarsubquery()
aboveavg = select(Experiment).where(
Experiment.accuracy > avgsubquery
)
topexperiments = session.exec(aboveavg).all()
# OR conditions
statement = select(Experiment).where(
or(
Experiment.algorithm == "ResNet50",
Experiment.algorithm == "EfficientNet",
)
)
# LIKE search
statement = select(Experiment).where(
Experiment.name.like("%transfer%")
)
# IN clause
algorithms = ["ResNet50", "VGG16", "EfficientNet"]
statement = select(Experiment).where(
Experiment.algorithm.in(algorithms)
)
# ORDER BY multiple columns
statement = (
select(Experiment)
.orderby(Experiment.accuracy.desc(), Experiment.createdat.desc())
)
return results
Practical Example: ML Experiment Tracking and Model Registry System
Here is a complete production-ready tracking system:
from fastapi import FastAPI, HTTPException, Depends, Query, UploadFile, File
from sqlmodel import Session, select, createengine, SQLModel, Field, Relationship
from typing import Optional
from datetime import datetime
import json
import os
Models
class MLModelBase(SQLModel):
name: str = Field(index=True)
version: str = Field(default="1.0.0")
framework: str
description: Optional[str] = None
class MLModelCreate(MLModelBase):
pass
class MLModelDB(MLModelBase, table=True):
tablename = "modelregistry"
id: Optional[int] = Field(default=None, primarykey=True)
modelpath: Optional[str] = None
filesizemb: Optional[float] = None
accuracy: Optional[float] = None
isactive: bool = Field(default=True, index=True)
isproduction: bool = Field(default=False)
createdat: datetime = Field(defaultfactory=datetime.utcnow)
metadatajson: Optional[str] = None
class MLModelRead(MLModelBase):
id: int
modelpath: Optional[str]
accuracy: Optional[float]
isactive: bool
isproduction: bool
createdat: datetime
Application
DATABASEURL = os.getenv("DATABASEURL", "sqlite:///./mlregistry.db")
engine = createengine(DATABASEURL)
app = FastAPI(title="ML Model Registry & Experiment Tracker")
def getsession():
with Session(engine) as session:
yield session
@app.onevent("startup")
def startup():
SQLModel.metadata.createall(engine)
@app.post("/models/register", responsemodel=MLModelRead)
def registermodel(
modeldata: MLModelCreate,
session: Session = Depends(getsession),
):
"""Register a new model to the registry"""
dbmodel = MLModelDB.modelvalidate(modeldata)
session.add(dbmodel)
session.commit()
session.refresh(dbmodel)
return dbmodel
@app.post("/models/{modelid}/promote")
def promotetoproduction(
modelid: int,
session: Session = Depends(getsession),
):
"""Promote a model to production"""
model = session.get(MLModelDB, modelid)
if not model:
raise HTTPException(statuscode=404, detail="Model not found")
# Demote current production model (if any)
currentprod = session.exec(
select(MLModelDB).where(
MLModelDB.name == model.name,
MLModelDB.isproduction == True,
)
).all()
for m in currentprod:
m.isproduction = False
session.add(m)
model.isproduction = True
session.add(model)
session.commit()
return {"message": f"Model {model.name} v{model.version} promoted to production"}
@app.get("/models/production/{modelname}", responsemodel=MLModelRead)
def getproductionmodel(
modelname: str,
session: Session = Depends(getsession),
):
"""Get the current production model"""
statement = select(MLModelDB).where(
MLModelDB.name == modelname,
MLModelDB.isproduction == True,
)
model = session.exec(statement).first()
if not model:
raise HTTPException(statuscode=404, detail="No production model found")
return model
@app.get("/models/compare")
def comparemodels(
modelids: str = Query(description="Comma-separated model IDs"),
session: Session = Depends(getsession),
):
"""Compare multiple models"""
ids = [int(id.strip()) for id in modelids.split(",")]
models = []
for modelid in ids:
model = session.get(MLModelDB, modelid)
if model:
models.append({
"id": model.id,
"name": model.name,
"version": model.version,
"accuracy": model.accuracy,
"framework": model.framework,
"isproduction": model.isproduction,
})
return {"models": models}
Tips and Best Practices
Base, Create, Read, and Update schemas for each model.ge, le, minlength, maxlength for model-level validation.index=True on fields that are frequently used in WHERE clauses.with Session) to avoid connection leaks.createall() in production. Always use Alembic for migrations.Conclusion
SQLModel is a modern ORM that is perfect for Python AI/ML applications. By combining the power of SQLAlchemy and Pydantic, SQLModel delivers a type-safe, validated, and seamlessly FastAPI-integrated development experience.
For AI projects, SQLModel is ideal for building backends for experiment tracking, model registries, dataset management, and various other data persistence needs. Its ease of use does not sacrifice power, as full access to SQLAlchemy remains available when needed.
Start with a small project using SQLite, then migrate to PostgreSQL when ready for production. With Alembic for migrations and FastAPI for the API layer, you have a solid stack for building a scalable AI platform.