Flyte Tutorial: Workflow Orchestration for Machine Learning and Data Engineering
Flyte is an open-source workflow orchestration platform purpose-built for machine learning, data engineering, and data science workloads. Developed by Union.ai and battle-tested at companies like Spotify, Lyft, and Freenome, Flyte provides a scalable, reproducible, and production-ready way to manage data and ML pipelines.
Unlike traditional workflow orchestrators such as Airflow that focus primarily on task scheduling, Flyte was designed from the ground up with deep understanding of ML requirements: data and model versioning, type safety, computation caching, and native container execution support. In this tutorial, we will learn Flyte from installation through building production-ready ML pipelines.
Why Flyte?
Before diving in, let's understand why Flyte is becoming an increasingly popular choice in the MLOps ecosystem:
Installation and Setup
Prerequisites
Ensure your system has:
- Python 3.8 or newer
- Docker (for running Flyte locally)
- pip or conda for package management
Installing Flytekit
Flytekit is the Python SDK for writing Flyte workflows:
pip install flytekit
For additional features like Pandas integration and visualization:
pip install flytekit[pandas]
pip install flytekitplugins-deck-standard
Running Flyte Locally
Flyte provides a sandbox for local development using Docker:
pip install flytectl
flytectl demo start
This command will run a complete Flyte cluster in Docker, including:
- Flyte Admin (API server)
- Flyte Console (web UI)
- Flyte Propeller (workflow engine)
- MinIO (object storage)
- PostgreSQL (metadata store)
Once successful, access the Flyte Console at http://localhost:30080/console.
For a lighter alternative, you can run workflows locally without a cluster:
pyflyte run myworkflow.py myworkflow --inputparam value
Core Concepts
Tasks
A task is the smallest unit of work in Flyte. It's a Python function decorated with @task:
from flytekit import task
@task
def sayhello(name: str) -> str:
return f"Hello, {name}!"
Note that type hints are mandatory. Flyte uses them for validation and data serialization.
Workflows
Workflows combine multiple tasks into a pipeline:
from flytekit import task, workflow
@task
def sayhello(name: str) -> str:
return f"Hello, {name}!"
@task
def greetinglength(greeting: str) -> int:
return len(greeting)
@workflow
def myworkflow(name: str) -> int:
greeting = sayhello(name=name)
length = greetinglength(greeting=greeting)
return length
Running Workflows
There are several ways to run workflows:
# Run locally (for development)
pyflyte run myworkflow.py myworkflow --name "Flyte"
Run on a Flyte cluster
pyflyte run --remote myworkflow.py myworkflow --name "Flyte"
Or run directly from Python:
if name == "main":
result = myworkflow(name="Flyte")
print(f"Result: {result}")
Basic Usage: Data Processing Pipeline
Let's build a simple data processing pipeline using Flyte:
import pandas as pd
from flytekit import task, workflow
from flytekit.types.file import FlyteFile
@task
def loaddata(filepath: str) -> pd.DataFrame:
"""Read dataset from a CSV file."""
df = pd.readcsv(filepath)
print(f"Loaded {len(df)} rows")
return df
@task
def cleandata(df: pd.DataFrame) -> pd.DataFrame:
"""Clean data: remove duplicates and missing values."""
initialrows = len(df)
df = df.dropduplicates()
df = df.dropna()
cleanedrows = len(df)
print(f"Cleaned: {initialrows} -> {cleanedrows} rows")
return df
@task
def computestatistics(df: pd.DataFrame) -> dict:
"""Calculate basic statistics from the dataset."""
stats = {
"rowcount": len(df),
"columncount": len(df.columns),
"columns": list(df.columns),
"numericmeans": df.selectdtypes(include="number").mean().todict(),
}
return stats
@workflow
def datapipeline(filepath: str) -> dict:
"""Complete pipeline: load, clean, and compute statistics."""
rawdata = loaddata(filepath=filepath)
clean = cleandata(df=rawdata)
stats = computestatistics(df=clean)
return stats
Run the pipeline:
pyflyte run datapipeline.py datapipeline --filepath "./data/sales.csv"
Advanced Usage: ML Training Pipeline
Model Training Pipeline with Caching
One of Flyte's best features is caching. By adding cache=True, tasks whose inputs haven't changed won't re-execute:
import pandas as pd
import numpy as np
from flytekit import task, workflow, Resources
from flytekit.types.file import FlyteFile
from sklearn.modelselection import traintestsplit
from sklearn.ensemble import RandomForestClassifier
from sklearn.metrics import accuracyscore, f1score
import joblib
import os
from dataclasses import dataclass
from mashumaro.mixins.json import DataClassJSONMixin
@dataclass
class TrainingConfig(DataClassJSONMixin):
testsize: float = 0.2
randomstate: int = 42
nestimators: int = 100
maxdepth: int = 10
@dataclass
class ModelMetrics(DataClassJSONMixin):
accuracy: float
f1score: float
featureimportances: dict
@task(cache=True, cacheversion="1.0")
def preparedataset(
filepath: str, config: TrainingConfig
) -> tuple[pd.DataFrame, pd.DataFrame, pd.Series, pd.Series]:
"""Load and split dataset into train/test."""
df = pd.readcsv(filepath)
X = df.drop(columns=["target"])
y = df["target"]
Xtrain, Xtest, ytrain, ytest = traintestsplit(
X, y, testsize=config.testsize, randomstate=config.randomstate
)
return Xtrain, Xtest, ytrain, ytest
@task(
requests=Resources(cpu="2", mem="4Gi"),
limits=Resources(cpu="4", mem="8Gi"),
)
def trainmodel(
Xtrain: pd.DataFrame,
ytrain: pd.Series,
config: TrainingConfig,
) -> FlyteFile:
"""Train a RandomForest model and save as a file."""
model = RandomForestClassifier(
nestimators=config.nestimators,
maxdepth=config.maxdepth,
randomstate=config.randomstate,
njobs=-1,
)
model.fit(Xtrain, ytrain)
modelpath = "/tmp/model.joblib"
joblib.dump(model, modelpath)
return FlyteFile(path=modelpath)
@task
def evaluatemodel(
modelfile: FlyteFile,
Xtest: pd.DataFrame,
ytest: pd.Series,
featurenames: list[str],
) -> ModelMetrics:
"""Evaluate the model and return metrics."""
modelfile.download()
model = joblib.load(modelfile.path)
predictions = model.predict(Xtest)
importances = dict(
zip(featurenames, model.featureimportances.tolist())
)
return ModelMetrics(
accuracy=accuracyscore(ytest, predictions),
f1score=f1score(ytest, predictions, average="weighted"),
featureimportances=importances,
)
@workflow
def mltrainingpipeline(
filepath: str, config: TrainingConfig = TrainingConfig()
) -> ModelMetrics:
"""Complete ML pipeline: prepare, train, evaluate."""
Xtrain, Xtest, ytrain, ytest = preparedataset(
filepath=filepath, config=config
)
modelfile = trainmodel(
Xtrain=Xtrain, ytrain=ytrain, config=config
)
metrics = evaluatemodel(
modelfile=modelfile,
Xtest=Xtest,
ytest=ytest,
featurenames=list(Xtrain.columns),
)
return metrics
Dynamic Workflows
Flyte supports dynamic workflows for cases where the workflow structure is determined at runtime:
from flytekit import task, workflow, dynamic
@task
def trainsinglemodel(
Xtrain: pd.DataFrame,
ytrain: pd.Series,
nestimators: int,
) -> float:
"""Train a single model and return accuracy."""
model = RandomForestClassifier(nestimators=nestimators, randomstate=42)
model.fit(Xtrain, ytrain)
return model.score(Xtrain, ytrain)
@dynamic
def hyperparametersearch(
Xtrain: pd.DataFrame,
ytrain: pd.Series,
nestimatorsoptions: list[int],
) -> list[float]:
"""Run training for each hyperparameter combination."""
results = []
for nest in nestimatorsoptions:
score = trainsinglemodel(
Xtrain=Xtrain,
ytrain=ytrain,
nestimators=nest,
)
results.append(score)
return results
@workflow
def searchpipeline(
filepath: str,
) -> list[float]:
Xtrain, Xtest, ytrain, ytest = preparedataset(
filepath=filepath, config=TrainingConfig()
)
scores = hyperparametersearch(
Xtrain=Xtrain,
ytrain=ytrain,
nestimatorsoptions=[50, 100, 200, 500],
)
return scores
Conditional Workflows
Flyte supports branching based on conditions:
from flytekit import task, workflow, conditional
@task
def calculateaccuracy(modelfile: FlyteFile, Xtest: pd.DataFrame, ytest: pd.Series) -> float:
modelfile.download()
model = joblib.load(modelfile.path)
return model.score(Xtest, ytest)
@task
def deploymodel(modelfile: FlyteFile) -> str:
return f"Model deployed: {modelfile.path}"
@task
def retrainnotification() -> str:
return "Model accuracy below threshold. Retraining needed."
@workflow
def deploymentpipeline(
filepath: str, accuracythreshold: float = 0.85
) -> str:
Xtrain, Xtest, ytrain, ytest = preparedataset(
filepath=filepath, config=TrainingConfig()
)
modelfile = trainmodel(
Xtrain=Xtrain, ytrain=ytrain, config=TrainingConfig()
)
accuracy = calculateaccuracy(
modelfile=modelfile, Xtest=Xtest, ytest=ytest
)
result = (
conditional("deploymentcheck")
.if(accuracy >= accuracythreshold)
.then(deploymodel(modelfile=modelfile))
.else()
.then(retrainnotification())
)
return result
Map Tasks for Parallel Execution
To run the same task with different inputs in parallel:
from flytekit import task, workflow, maptask
@task
def process
partition(partitionid: int) -> dict:
"""Process a single data partition."""
return {
"partition
id": partitionid,
"rows
processed": partitionid 1000,
"status": "completed",
}
@task
def aggregate
results(results: list[dict]) -> dict:
"""Combine results from all partitions."""
totalrows = sum(r["rowsprocessed"] for r in results)
return {
"totalpartitions": len(results),
"total
rows": totalrows,
}
@workflow
def parallel
pipeline(numpartitions: int = 10) -> dict:
partition
ids = list(range(numpartitions))
results = map
task(processpartition)(partitionid=partitionids)
summary = aggregate
results(results=results)
return summary
Integrating with the ML Ecosystem
Flyte + MLflow
Flyte can be integrated with MLflow for experiment tracking:
import mlflow
from flytekit import task, workflow
from flytekitplugins.mlflow import mlflowautolog
@task(enabledeck=True)
@mlflowautolog(framework=mlflow.sklearn)
def trainwithtracking(
Xtrain: pd.DataFrame,
ytrain: pd.Series,
nestimators: int,
) -> FlyteFile:
"""Training with MLflow autologging."""
model = RandomForestClassifier(nestimators=nestimators)
model.fit(Xtrain, ytrain)
modelpath = "/tmp/model.joblib"
joblib.dump(model, modelpath)
return FlyteFile(path=modelpath)
Flyte + Great Expectations
Validate data using Great Expectations within a Flyte pipeline:
from flytekit import task
import greatexpectations as gx
@task
def validatedata(df: pd.DataFrame) -> bool:
"""Validate data quality before training."""
context = gx.getcontext()
datasource = context.datasources.addpandas("pandasds")
dataasset = datasource.adddataframeasset("myasset")
batch = dataasset.addbatchdefinitionwholedataframe("batch").getbatch(
batchparameters={"dataframe": df}
)
expectationsuite = context.suites.add(
gx.ExpectationSuite(name="trainingdatasuite")
)
expectationsuite.addexpectation(
gx.expectations.ExpectColumnValuesToNotBeNull(column="target")
)
expectationsuite.addexpectation(
gx.expectations.ExpectColumnValuesToBeBetween(
column="feature1", minvalue=0, maxvalue=100
)
)
validationresult = batch.validate(expectationsuite)
return validationresult.success
Container Tasks
To run tasks in custom containers (e.g., GPU training):
from flytekit import task, ImageSpec
customimage = ImageSpec(
name="ml-training",
packages=[
"torch==2.1.0",
"transformers==4.36.0",
"scikit-learn==1.3.0",
],
pythonversion="3.11",
cuda="12.1",
registry="ghcr.io/myorg",
)
@task(containerimage=customimage)
def traindeeplearningmodel(
datasetpath: str,
epochs: int = 10,
learningrate: float = 1e-4,
) -> FlyteFile:
"""Train a deep learning model in a GPU container."""
import torch
from transformers import AutoModelForSequenceClassification, Trainer
model = AutoModelForSequenceClassification.frompretrained(
"bert-base-uncased", numlabels=2
)
modelpath = "/tmp/dlmodel"
model.savepretrained(modelpath)
return FlyteFile(path=modelpath)
Scheduling and Monitoring
Launch Plans
Launch Plans let you schedule workflows:
from flytekit import LaunchPlan, CronSchedule
dailytraining = LaunchPlan.getorcreate(
name="dailytrainingpipeline",
workflow=mltrainingpipeline,
schedule=CronSchedule(schedule="0 2 "),
defaultinputs={
"filepath": "/data/dailydataset.csv",
"config": TrainingConfig(nestimators=200),
},
)
Notifications
Configure notifications when workflows complete or fail:
from flytekit import LaunchPlan, Email
monitoredpipeline = LaunchPlan.getorcreate(
name="monitoredpipeline",
workflow=mltrainingpipeline,
notifications=[
Email(
phases=[
WorkflowExecutionPhase.FAILED,
WorkflowExecutionPhase.SUCCEEDED,
],
recipientsemail=["team@company.com"],
)
],
)
Best Practices
1. Use Type Hints Consistently
Always explicitly define input and output types. This helps Flyte with validation and serialization:
@task
def goodtask(data: pd.DataFrame, threshold: float) -> list[str]:
return [col for col in data.columns if data[col].mean() > threshold]
2. Leverage Caching for Expensive Tasks
Add cache=True to tasks that take a long time and produce deterministic results:
@task(cache=True, cacheversion="1.0")
def expensive
computation(data: pd.DataFrame) -> pd.DataFrame:
return data.apply(complextransformation)
Update cacheversion when the task logic changes so old cache entries aren't reused.
3. Define Resource Requirements
Specify resource needs for each task so Kubernetes can allocate appropriately:
@task(
requests=Resources(cpu="2", mem="4Gi", gpu="1"),
limits=Resources(cpu="4", mem="8Gi", gpu="1"),
)
def gputrainingtask(data: pd.DataFrame) -> FlyteFile:
pass
4. Use Dataclasses for Structured Data
For complex configurations and results, use dataclasses:
from dataclasses import dataclass
from mashumaro.mixins.json import DataClassJSONMixin
@dataclass
class ExperimentResult(DataClassJSONMixin):
modelname: str
accuracy: float
parameters: dict
trainingdurationseconds: float
5. Test Locally Before Deploying
Always test workflows locally before deploying to a cluster:
if name == "main":
result = mltrainingpipeline(
filepath="./testdata.csv",
config=TrainingConfig(nestimators=10),
)
print(f"Metrics: {result}")
6. Use ImageSpec for Dependency Management
Define dependencies at the task level, not the project level:
dataimage = ImageSpec(
name="data-processing",
packages=["pandas==2.1.0", "numpy==1.26.0"],
)
ml
image = ImageSpec(
name="ml-training",
packages=["scikit-learn==1.3.0", "xgboost==2.0.0"],
)
@task(containerimage=dataimage)
def processdata(raw: pd.DataFrame) -> pd.DataFrame:
pass
@task(containerimage=mlimage)
def train(data: pd.DataFrame) -> FlyteFile:
pass
7. Structured Error Handling
Use Flyte's retry mechanism instead of manual try-catch:
@task(retries=3, timeout=timedelta(hours=1))
def flakyapicall(endpoint: str) -> dict:
"""Task with automatic retry on failure."""
import requests
response = requests.get(endpoint, timeout=30)
response.raiseforstatus()
return response.json()
Deploying to Production
Registering Workflows to a Cluster
# Build and push container image
pyflyte --pkgs myproject package --image ghcr.io/myorg/my-project:latest
Register to cluster
flytectl register files --project my-project --domain production --archive flyte-package.tgz
Recommended Project Structure
my-flyte-project/
├── workflows/
│ ├── init.py
│ ├── datapipeline.py
│ ├── trainingpipeline.py
│ └── deploymentpipeline.py
├── tasks/
│ ├── init.py
│ ├── datatasks.py
│ ├── modeltasks.py
│ └── evaluationtasks.py
├── configs/
│ ├── trainingconfig.py
│ └── deploymentconfig.py
├── tests/
│ ├── testdatatasks.py
│ └── testtrainingpipeline.py
├── Dockerfile
├── requirements.txt
└── pyproject.toml
Conclusion
Flyte is a workflow orchestration platform purpose-built for the needs of modern machine learning and data engineering. With type safety, caching, versioning, and native Kubernetes support, Flyte helps data teams build reproducible and scalable pipelines.
Key advantages of Flyte over alternatives:
- Type safety catches errors that would otherwise surface only at runtime
- Built-in caching saves time and compute costs
- Automatic versioning ensures reproducibility
- Dynamic workflows provide flexibility for complex use cases
- Multi-tenancy simplifies cross-team collaboration
To get started, build a simple pipeline with two or three tasks, run it locally, then deploy to a Flyte sandbox cluster. As your needs grow, leverage advanced features like map tasks, conditional workflows, and integrations with other ML tools.
Complete documentation is available on the official Flyte website, and an active Slack community is ready to help if you run into issues.