What you will be able to do
- Recognise when a chain needs a custom pyfunc model rather than autologging or a built-in flavor
- Split a chain's logic correctly between load_context and predict
- Write a PythonModel that formats inputs before the model call and post-processes outputs after it
- Map the steps of a RAG chain onto the preprocessing and postprocessing hooks of a pyfunc model
Key concept
The PythonModel contract (load_context vs predict) — A custom pyfunc model is a subclass of mlflow.pyfunc.PythonModel. Anything expensive that only has to be loaded once, like weights, tokenizers or clients, goes in load_context. Everything that has to run on each request, including preprocessing, the model call and postprocessing, goes in predict.
1.Why a chain becomes a pyfunc model
MLflow packages models in *flavors* that serving and inference platforms understand. Examples are python-function, pytorch and sklearn. A built-in flavor such as sklearn can only serve what the library's own predict returns. A generative AI application usually needs more than that. It cleans up or rewrites the user's request, calls one or more models, and then shapes the raw result into a form the caller can use. MLflow's Python function flavor, pyfunc, is built for this, because it "provides flexibility to deploy any piece of Python code or any Python model."
Databricks lists the situations where you would write custom pyfunc code:
- The model needs preprocessing before inputs can reach its predict function.
- MLflow doesn't natively support the model framework.
- The model's raw outputs have to be post-processed before anyone can use them.
- The model has per-request branching logic.
- You want to deploy fully custom code as a model.
The first and third items are what this objective is about. When a chain has a preprocessing stage, a model call and a postprocessing stage, all three travel together as one pyfunc model. A Model Serving endpoint then serves that model, which Databricks calls a *custom model*.
Checkpoint 1 of 4· Check yourself
A chain's model returns raw scores that the calling application cannot consume directly, so they must be converted before being returned. Which reason from the Databricks list of pyfunc scenarios applies?
Raw outputs that must be reshaped before the caller can use them are the documented post-processing scenario for pyfunc. Nothing here suggests an unsupported framework.
“Your application requires the model's raw outputs to be post-processed for consumption.”Source: docs.databricks.com
Sources1
2.The two methods: load_context and predict
To write a pyfunc model, you subclass mlflow.pyfunc.PythonModel. The Databricks guide names two functions to implement when you package arbitrary Python code this way.
**load_context(self, context)** runs once, when the model loads. Put anything here that only has to be loaded once for the model to work, such as model weights, a tokenizer, or helper objects built from packaged files. The reason is speed. Loading these things up front keeps the number of artifacts loaded during predict to a minimum, and that makes inference faster.
**predict(self, context, model_input)** "houses all the logic that is run every time an input request is made." Preprocessing, the call to the underlying model and postprocessing all happen here, on every request.
The rule is: one-time setup goes in load_context, and per-request work goes in predict. Suppose a chain reloads its weights or rebuilds its tokenizer inside predict. It still returns correct answers, but it repeats the slowest work on every call, which is exactly what load_context exists to avoid.
Checkpoint 2 of 4· Check yourself
A pyfunc chain reads a tokenizer from its packaged artifacts and then normalises each incoming query with it. Where should each piece go?
Loading the tokenizer is a one-time job, so it belongs in load_context. Normalising a query depends on that request, so it has to run in predict.
“predict - this function houses all the logic that is run every time an input request is made.”Source: docs.databricks.com
Sources1
3.Writing the pre- and post-processing hooks
The Databricks example below shows the standard shape. Preprocessing and postprocessing each live in their own helper method, and predict calls them in order: format the input, call the model, format the output. Look at what load_context does too. It loads the weights and builds a tokenizer from paths in context.artifacts. It also imports a tokenizer class from a shared preprocessing_utils module that was packaged with the model.
class CustomModel(mlflow.pyfunc.PythonModel):
def load_context(self, context):
self.model = torch.load(context.artifacts["model-weights"])
from preprocessing_utils.my_custom_tokenizer import CustomTokenizer
self.tokenizer = CustomTokenizer(context.artifacts["tokenizer_cache"])
def format_inputs(self, model_input):
# insert some code that formats your inputs
pass
def format_outputs(self, outputs):
predictions = (torch.sigmoid(outputs)).data.numpy()
return predictions
def predict(self, context, model_input):
model_input = self.format_inputs(model_input)
outputs = self.model.predict(model_input)
return self.format_outputs(outputs)format_outputs applies a sigmoid to the raw outputs and converts the result to a NumPy array. The caller gets usable predictions instead of raw model output. This matches the documented case where raw outputs have to be post-processed before anyone can use them.
Checkpoint 3 of 4· Exam question
A generative AI engineer is packaging a RAG chain as a custom MLflow pyfunc model by subclassing `mlflow.pyfunc.PythonModel`. The chain needs to load a large embedding lookup file once and reuse it across every request, while the retrieval-formatting and response-cleanup logic must run on each incoming request. Where should the engineer place the one-time loading code versus the per-request logic?
Correct answer: A — Load the embedding lookup file inside `load_context`, and place the retrieval-formatting and response-cleanup logic inside `predict`.
- A. This is correct: `load_context` runs once when the model is initialized, making it the right place for expensive, reusable setup like loading a large lookup file, while `predict` runs on every request and should hold per-request logic like formatting and cleanup.
- B. This inverts the responsibilities: reloading a large lookup file on every call inside `predict` adds unnecessary latency to each request, and `load_context` is not the place for per-request formatting logic since it only runs once at initialization.
- C. `load_context` runs in production model serving as well as local testing, not just locally, and reloading the lookup file on every `predict` call wastes time and resources that one-time initialization would avoid.
- D. `predict` is where inference and request-specific logic belongs, not merely schema validation, and putting per-request formatting into `load_context` means it would only execute once and never process actual request data.
Sources1
4.Where pre- and post-processing sit in a RAG chain
Databricks uses the term *RAG chain* for the series of steps that runs at inference time when a user submits a request. Two of those steps line up directly with the hooks you just saw.
1. (Optional) User query preprocessing. The query is reshaped so it works better for searching the vector database. That can mean putting it into a template, using another model to rewrite the request, or extracting keywords. The output is a *retrieval query*. 2. Retrieval. The retrieval query is turned into an embedding with the same embedding model that embedded the document chunks. The most similar chunks come back. 3. Prompt augmentation. The retrieved context is combined with the user's query in a prompt template. 4. LLM generation. The LLM produces a response grounded in that context. 5. (Optional) Post-processing. The response "may be processed further to apply additional business logic, add citations, or otherwise refine the generated text."
In a pyfunc chain, step 1 is your input-formatting code and step 5 is your output-formatting code. Steps 2 to 4 are what predict does between the two. Guardrails can also be added anywhere along the chain, for example filtering requests or checking user permissions before data sources are accessed. Those also fit naturally inside predict.
Checkpoint 4 of 4· Put it in order
Put the steps of a RAG chain at inference time in order.
- 1.Post-process the response, for example by adding citations
- 2.Embed the retrieval query and retrieve the most similar chunks
- 3.Preprocess the user query into a retrieval query
- 4.Augment the prompt with the retrieved context
- 5.Generate a response with the LLM
Preprocessing produces the retrieval query that retrieval needs. Retrieval supplies the context for the augmented prompt. The LLM consumes that prompt, and postprocessing refines what it generates.
“The LLM takes the augmented prompt, which includes the user's query and retrieved supporting data, as input.”Source: docs.databricks.com
Sources2
Exam traps
Each one states something that sounds right. Open it to see what is actually true.
1.Loading weights or tokenizers inside predict is fine, because the answers come out the same.Why is that wrong?
predict runs on every request. One-time loading belongs in load_context so that fewer artifacts are loaded during predict and inference stays fast.
Covered in The two methods: load_context and predict
Sources
Every claim above is drawn from one of these pages, quoted as it was written on the date shown.
- 1.https://docs.databricks.com/aws/en/machine-learning/model-serving/deploy-custom-python-codeOfficial docs
“MLflow's Python function, pyfunc, provides flexibility to deploy any piece of Python code or any Python model.”
↩︎ Why a chain becomes a pyfunc model“Your application requires the model's raw outputs to be post-processed for consumption.”
↩︎ Why a chain becomes a pyfunc model“Your model requires preprocessing before inputs can be passed to the model's predict function.”
↩︎ Why a chain becomes a pyfunc model“predict - this function houses all the logic that is run every time an input request is made.”
↩︎ The two methods: load_context and predict“This is critical so that the system minimize the number of artifacts loaded during the predict function, which speeds up inference.”
↩︎ The two methods: load_context and predict“Your application requires the model's raw outputs to be post-processed for consumption.”
↩︎ Writing the pre- and post-processing hooks“load_context - anything that needs to be loaded just one time for the model to operate should be defined in this function.”
↩︎ Key concept“This is critical so that the system minimize the number of artifacts loaded during the predict function, which speeds up inference.”
↩︎ Exam trap 1 - 2.https://docs.databricks.com/aws/en/agents/tutorials/ai-cookbook/fundamentals-inference-chain-ragOfficial docs
“The series, or chain of steps that are invoked at inference time is commonly referred to as the RAG chain.”
↩︎ Where pre- and post-processing sit in a RAG chain“This can involve formatting the query within a template, using another model to rewrite the request, or extracting keywords to aid retrieval.”
↩︎ Where pre- and post-processing sit in a RAG chain“The LLM's response may be processed further to apply additional business logic, add citations, or otherwise refine the generated text”
↩︎ Where pre- and post-processing sit in a RAG chain“The LLM takes the augmented prompt, which includes the user's query and retrieved supporting data, as input.”
↩︎ Checkpoint
Also cited
“Use this if you have a custom model or if you need extra steps before or after inference.”
↩︎ Prediction