trainhub-api / app /models.py
Taimwe's picture
Rewrite ORM models (simplify file model)
db913fc verified
Raw History Blame Contribute Delete
3.69 kB
"""SQLAlchemy ORM models."""
import datetime
import uuid
from sqlalchemy import Column, String, Integer, Text, Float, DateTime, JSON
from .database import Base
def _id() -> str:
return uuid.uuid4().hex
class User(Base):
__tablename__ = "users"
id = Column(String, primary_key=True, default=_id)
username = Column(String, unique=True, index=True, nullable=False)
email = Column(String, unique=True, index=True, nullable=False)
hashed_password = Column(String, nullable=False)
created_at = Column(DateTime, default=datetime.datetime.utcnow)
class ModelRepo(Base):
__tablename__ = "model_repos"
id = Column(String, primary_key=True, default=_id)
name = Column(String, nullable=False)
owner = Column(String, index=True, nullable=False)
description = Column(Text, default="")
pipeline_tag = Column(String, default="text-generation")
tags = Column(JSON, default=list)
license = Column(String, default="apache-2.0")
hf_repo_id = Column(String, nullable=True)
downloads = Column(Integer, default=0)
likes = Column(Integer, default=0)
created_at = Column(DateTime, default=datetime.datetime.utcnow)
class DatasetRepo(Base):
__tablename__ = "dataset_repos"
id = Column(String, primary_key=True, default=_id)
name = Column(String, nullable=False)
owner = Column(String, index=True, nullable=False)
description = Column(Text, default="")
tags = Column(JSON, default=list)
license = Column(String, default="apache-2.0")
columns = Column(JSON, default=list) # schema: list of column names
num_rows = Column(Integer, default=0)
size_bytes = Column(Integer, default=0)
downloads = Column(Integer, default=0)
created_at = Column(DateTime, default=datetime.datetime.utcnow)
class SpaceRepo(Base):
__tablename__ = "space_repos"
id = Column(String, primary_key=True, default=_id)
name = Column(String, nullable=False)
owner = Column(String, index=True, nullable=False)
description = Column(Text, default="")
sdk = Column(String, default="gradio")
status = Column(String, default="not-built") # not-built | building | running | sleeping
likes = Column(Integer, default=0)
created_at = Column(DateTime, default=datetime.datetime.utcnow)
class TrainingJob(Base):
__tablename__ = "training_jobs"
id = Column(String, primary_key=True, default=_id)
name = Column(String, nullable=False)
owner = Column(String, index=True, nullable=False)
base_model = Column(String, nullable=False)
dataset = Column(String, nullable=False)
job_type = Column(String, default="scratch") # scratch | finetune
status = Column(String, default="queued") # queued | running | completed | failed | cancelled
config = Column(JSON, default=dict)
metrics = Column(JSON, default=dict)
total_epochs = Column(Integer, default=0)
current_epoch = Column(Integer, default=0)
loss = Column(Float, nullable=True)
error = Column(Text, nullable=True)
created_at = Column(DateTime, default=datetime.datetime.utcnow)
started_at = Column(DateTime, nullable=True)
finished_at = Column(DateTime, nullable=True)
class StoredFile(Base):
__tablename__ = "stored_files"
id = Column(String, primary_key=True, default=_id)
repo_type = Column(String, index=True, nullable=False) # model | dataset
repo_id = Column(String, index=True, nullable=False) # owning repo id (polymorphic)
filename = Column(String, nullable=False)
size_bytes = Column(Integer, default=0)
sha256 = Column(String, nullable=True)
uploaded_at = Column(DateTime, default=datetime.datetime.utcnow)