Aller au contenu principal

Affinement supervisé (complet) et serving de Qwen3.5-0.8B

Ouvrir dans Databricks

Ajuster finement le modèle compact open-weight Qwen3.5-0.8B-Base grand modèle de langage sur AI Runtime (GPU Serverless), puis déployez-le derrière un Model Serving Endpoint. Cet exemple s’exécute de bout en bout sur un unique GPU H100 et vous montre comment :

  • Exécutez l' affinement supervisé (SFT) avec le SFTTrainerde TRL sur un dataset de suivi d'instructions
  • Comparez les réponses du modèle avant et après l'affinement pour voir l'effet du SFT
  • Enregistrez le modèle affiné dans Unity Catalog pour la gouvernance et le déploiement
  • Servir le modèle derrière un endpoint Custom Foundation Model exécutant un serveur compatible OpenAI vLLM

Concepts clés :

  • Affinement supervisé (SFT): poursuit l'entraînement d'un modèle de base sur des paires instruction/réponse sélectionnées afin qu'il suive les instructions dans le style cible
  • TRL: une bibliothèque pour l’affinement supervisé et l’apprentissage par renforcement des modèles de langage
  • Service de modèle de fondation personnalisé: dessert vos propres pondérations de LLM affinées sur Model Serving compatible GPU avec une API compatible OpenAI
remarque

Cet exemple requiert l'environnement AI Runtime version 6 ou supérieure (l'étape de diffusion utilise vLLM et flashinfer, qui sont intégrés dans la v6).

Se connecter au compute GPU serverless​

Ce notebook nécessite un compute GPU serverless. Pour vous connecter :

  1. Cliquez sur le sélecteur de compute du notebook en haut à droite et sélectionnez Serverless GPU .
  2. Sur la droite, cliquez sur le bouton de l’environnement.
  3. Sélectionnez H100 comme Accélérateur .
  4. Choisissez AI v6 dans l'environnement de base.
  5. Cliquez sur Appliquer .

Configuration​

La cellule suivante définit les widgets pour l'emplacement dans Unity Catalog où le modèle affiné est enregistré. Le modèle est enregistré en tant que {uc_catalog}.{uc_schema}.{uc_model_name}, et l'endpoint de mise en service est nommé {uc_model_name}-endpoint.

Définit deploy_endpoint sur false pour s'arrêter après l'enregistrement du modèle et le test vLLM local, sans déployer d'endpoint de service géré.

Python
dbutils.widgets.text("uc_catalog", "main")
dbutils.widgets.text("uc_schema", "default")
dbutils.widgets.text("uc_model_name", "qwen3_5_0_8b_sft")
dbutils.widgets.dropdown("deploy_endpoint", "true", ["true", "false"])

UC_CATALOG = dbutils.widgets.get("uc_catalog")
UC_SCHEMA = dbutils.widgets.get("uc_schema")
UC_MODEL_NAME_BASE = dbutils.widgets.get("uc_model_name")
# Whether to deploy the managed serving endpoint (Steps 8-9). Set to "false" to stop after
# registration and the local vLLM test (for example, on workspaces where entrypoint-based
# Custom Foundation Model serving is not enabled).
DEPLOY_ENDPOINT = dbutils.widgets.get("deploy_endpoint").lower() == "true"

print(f"UC_CATALOG: {UC_CATALOG}")
print(f"UC_SCHEMA: {UC_SCHEMA}")
print(f"UC_MODEL_NAME: {UC_MODEL_NAME_BASE}")
print(f"DEPLOY_ENDPOINT: {DEPLOY_ENDPOINT}")

Importer des bibliothèques​

Charger les bibliothèques utilisées dans l’ensemble du Notebook. L’environnement AI Runtime v6 inclut déjà torch, transformers, trl et datasets; aucune installation n’est donc requise.

Python
import torch
import pandas as pd
from datasets import load_dataset, Dataset
import transformers
from transformers import TrainingArguments, AutoTokenizer, AutoModelForCausalLM
from trl import SFTTrainer, SFTConfig

Étape 1 : Charger le modèle de base et le tokeniseur​

Charger le modèle pré-entraîné Qwen3.5-0.8B-Base checkpoint depuis le Hugging Face Hub.

  • Architecture : Qwen3 (transformateur orienté décodeur uniquement, ~0,8 milliard de parameters)
  • Point de contrôle « Base » : aucun instruction-tuning appliqué pour l'instant. Il s'agit du modèle à affiner ci-dessous.
  • Le modèle est déplacé vers le GPU immédiatement après le chargement pour accélérer l'inférence et l'entraînement.
Python
# Qwen3.5-0.8B-Base: ~0.8B parameter decoder-only model.
# "Base" = no instruction-tuning yet; this is the checkpoint fine-tuned below.
model_name = "Qwen/Qwen3.5-0.8B-Base"
tokenizer = AutoTokenizer.from_pretrained(model_name)
model = AutoModelForCausalLM.from_pretrained(model_name)
Python
# Move all model weights to the GPU for faster inference and training.
model.to("cuda")

Configurer le tokenizer​

Les checkpoints de base sont livrés sans template de chat. Définissez un template System / User / Assistant Jinja minimal afin que le tokenizer puisse formater correctement les prompts pour les conversations à un ou plusieurs tours. Définissez pad_token = eos_token car le vocabulaire de base n'a pas de token de remplissage dédié.

Python
# Base checkpoints ship without a chat template.
# Define a minimal System / User / Assistant Jinja template so the tokenizer
# can format both single-turn and multi-turn conversations correctly.
if not tokenizer.chat_template:
print("No chat template — applying default template")
tokenizer.chat_template = """{% for message in messages %}
{% if message['role'] == 'system' %}System: {{ message['content'] }}\n
{% elif message['role'] == 'user' %}User: {{ message['content'] }}\n
{% elif message['role'] == 'assistant' %}Assistant: {{ message['content'] }} <|endoftext|>
{% endif %}
{% endfor %}"""

# Set pad_token = eos_token because the base vocabulary has no dedicated pad token.
if not tokenizer.pad_token:
print("No pad token — using eos_token as pad_token")
tokenizer.pad_token = tokenizer.eos_token
Python
print(tokenizer.chat_template[0:100])
print(tokenizer.pad_token)

Étape 2 : Inférence de référence (pré-SFT)​

Exécutez une inférence de vérification rapide avec le modèle de base (non affiné) . La réponse sert ici de point de référence. Comparez-le avec la sortie du modèle SFT à l’étape 5.

Python
# Build a single-turn chat in OpenAI-style message format.
# The tokenizer's chat template will convert this list into a formatted prompt string.
messages = []
user_message = "Give me a one-sentence introduction to LLMs."
messages.append({"role": "user", "content": user_message})
messages
Python
# Render the message list into a raw text string using the chat template.
# add_generation_prompt=True appends the "Assistant:" prefix to trigger generation.
# enable_thinking=False disables Qwen3's chain-of-thought reasoning mode.
prompt = tokenizer.apply_chat_template(
messages,
tokenize=False,
add_generation_prompt=True,
enable_thinking=False,
)
prompt
Python
# Tokenize the prompt string and move tensors to the same device as the model.
inputs = tokenizer(prompt, return_tensors="pt").to(model.device)
inputs
Python
max_new_tokens = 100
with torch.no_grad(): # no gradient tracking needed during inference
outputs = model.generate(
**inputs, # pass tokenized prompt (input_ids + attention_mask)
max_new_tokens=max_new_tokens,
do_sample=False, # greedy decoding — deterministic output
pad_token_id=tokenizer.eos_token_id,
eos_token_id=tokenizer.eos_token_id,
)

outputs
Python
# Slice off the prompt tokens; decode only the newly generated portion.
input_len = inputs["input_ids"].shape[1]
generated_ids = outputs[0][input_len:]
response = tokenizer.decode(generated_ids, skip_special_tokens=True).strip()
response

Fonctions d'aide​

Utilitaires définis pour ce notebook :

  • generate_responses: formate un prompt avec le template de chat et exécute le décodage glouton
  • test_model_with_questions: évalue une liste de questions et affiche les résultats des modèles côte à côte
  • load_model_and_tokenizer: charge le modèle + le tokeniseur avec un placement GPU facultatif et l’application de patches au template
  • display_dataset: restitue les 3 premières lignes d'un dataset au format de chat sous forme de tableau lisible
Python
def generate_responses(model, tokenizer, user_message, system_message=None,
max_new_tokens=100):
# Format chat using tokenizer's chat template
messages = []
if system_message:
messages.append({"role": "system", "content": system_message})

# Assume the data are all single-turn conversations
messages.append({"role": "user", "content": user_message})

prompt = tokenizer.apply_chat_template(
messages,
tokenize=False,
add_generation_prompt=True,
enable_thinking=False,
)

inputs = tokenizer(prompt, return_tensors="pt").to(model.device)
# Recommended to use vllm, sglang or TensorRT
with torch.no_grad():
outputs = model.generate(
**inputs,
max_new_tokens=max_new_tokens,
do_sample=False,
pad_token_id=tokenizer.eos_token_id,
eos_token_id=tokenizer.eos_token_id,
)
input_len = inputs["input_ids"].shape[1]
generated_ids = outputs[0][input_len:]
response = tokenizer.decode(generated_ids, skip_special_tokens=True).strip()

return response
Python
def test_model_with_questions(model, tokenizer, questions,
system_message=None, title="Model Output"):
print(f"\n=== {title} ===")
for i, question in enumerate(questions, 1):
response = generate_responses(model, tokenizer, question,
system_message)
print(f"\nModel Input {i}:\n{question}\nModel Output {i}:\n{response}\n")
Python
def load_model_and_tokenizer(model_name, use_gpu = False):

# Load base model and tokenizer
tokenizer = AutoTokenizer.from_pretrained(model_name)
model = AutoModelForCausalLM.from_pretrained(model_name)

if use_gpu:
model.to("cuda")

if not tokenizer.chat_template:
tokenizer.chat_template = """{% for message in messages %}
{% if message['role'] == 'system' %}System: {{ message['content'] }}\n
{% elif message['role'] == 'user' %}User: {{ message['content'] }}\n
{% elif message['role'] == 'assistant' %}Assistant: {{ message['content'] }} <|endoftext|>
{% endif %}
{% endfor %}"""

# Tokenizer config
if not tokenizer.pad_token:
tokenizer.pad_token = tokenizer.eos_token

return model, tokenizer
Python
def display_dataset(dataset):
# Visualize the dataset
rows = []
for i in range(3):
example = dataset[i]
user_msg = next(m['content'] for m in example['messages']
if m['role'] == 'user')
assistant_msg = next(m['content'] for m in example['messages']
if m['role'] == 'assistant')
rows.append({
'User Prompt': user_msg,
'Assistant Response': assistant_msg
})

# Display as table
df = pd.DataFrame(rows)
pd.set_option('display.max_colwidth', None) # Avoid truncating long strings
display(df)

Étape 3 : Charger le dataset d'entraînement​

Charger banghua/DL-SFT-Dataset à partir du Hugging Face Hub, un dataset de suivi des instructions issu du cours SFT de DeepLearning.AI. Chaque exemple est une liste messages avec des tours user et assistant.

Cet exemple utilise un sous-ensemble de 100 exemples pour réduire le temps d'entraînement. Augmentez la taille du sous-ensemble pour un véritable affinement.

Python
train_dataset = load_dataset("banghua/DL-SFT-Dataset")['train']

train_dataset=train_dataset.select(range(100))

display_dataset(train_dataset)

Étape 4 : affinement SFT​

Utilisez SFTTrainer pour exécuter l’ affinement supervisé . Principaux choix d’hyperparamètres pour cette démonstration :

parameter

Valeur

Notes

learning_rate

8e-5

Point de départ standard pour l’affinement sur les petits modèles

num_train_epochs

1

Passe unique pour la démo ; augmentez pour l'entraînement réel

per_device_train_batch_size

1

Configurer avec gradient_accumulation_steps

gradient_accumulation_steps

8

Taille de batch effective = 1 × 8 = 8

gradient_checkpointing

False

Désactivé pour plus de rapidité ; activez cette option pour réduire la VRAM sur les grands modèles

parameter

Valeur

Notes

learning_rate

8e-5

Point de départ standard pour l’affinement sur les petits modèles

num_train_epochs

1

Passe unique pour la démo ; augmentez pour l'entraînement réel

per_device_train_batch_size

1

Configurer avec gradient_accumulation_steps

gradient_accumulation_steps

8

Taille de batch effective = 1 × 8 = 8

gradient_checkpointing

False

Désactivé pour plus de rapidité ; activez cette option pour réduire la VRAM sur les grands modèles

Python
# SFTConfig is a superset of HuggingFace TrainingArguments with SFT-specific defaults.
sft_config = SFTConfig(
# --- Training hyperparameters ---
learning_rate=8e-5, # standard starting point for SFT on small models
num_train_epochs=1, # single pass for demo; increase for real training
per_device_train_batch_size=1, # fits H100 VRAM; tune together with gradient_accumulation_steps
gradient_accumulation_steps=8, # effective batch size = 1 × 8 = 8
gradient_checkpointing=False, # disable for speed; enable to reduce VRAM on larger models
logging_steps=2,

# --- Logging ---
report_to=[], # disable W&B / MLflow / etc.
logging_strategy="steps",
logging_first_step=True,

# --- Checkpointing ---
output_dir="./checkpoints/tiny-finetune-exp1",
run_name="tiny-finetune-exp1-run", # must differ from output_dir to avoid W&B conflicts
save_strategy="no", # set to "epoch" to persist a final checkpoint
save_total_limit=1
)
Python
# SFTTrainer handles dataset formatting, tokenization, and response-label masking
# automatically based on the tokenizer's chat template.
sft_trainer = SFTTrainer(
model=model,
args=sft_config,
train_dataset=train_dataset,
processing_class=tokenizer # replaces the deprecated tokenizer= arg in TRL 0.12+
)

# Run training — loss should decrease over the single epoch on this 100-example subset.
sft_trainer.train()

Étape 5 : Évaluation post-SFT​

Comparez les réponses du modèle ajusté aux mêmes questions utilisées à l’étape 2. Recherchez un formatage amélioré et un style de respect des consignes qui reflète la distribution des données d’entraînement.

Python
questions = [
"Calculate 1+1-1",
"What's the difference between thread and process?"
]
test_model_with_questions(sft_trainer.model, tokenizer, questions,
title="Fine-tuned model output")

Étape 6 : Diffuser le modèle affiné avec vLLM​

Une fois l'entraînement terminé, déployez le modèle affiné derrière un endpoint de service Custom Foundation Model Databricks qui exécute un serveur compatible vLLM OpenAI.

Détail intéressant : Qwen/Qwen3.5-0.8B-Base est en réalité un modèle multimodal , et AutoModelForCausalLM ci-dessus n’a chargé que son socle textuel (Qwen3_5ForCausalLM). vLLM (env. AI v6) peut servir ce socle textuel en mode natif, mais quelques adaptations sont nécessaires car il s’agit de la partie textuelle d’un modèle vision-langage :

  • Enregistrez les pondérations ajustées finement avec le Qwen3_5Configcomposite (celui qui transporte vision_config) ainsi que les fichiers de processeur d’origine du modèle, sinon le processeur de vLLM rejette le point de contrôle.
  • Au lancement, indiquez à vLLM qu'il y a zéro image/vidéo (--limit-mm-per-prompt) afin qu'il n'interagisse jamais avec la tour de vision (absente), exécutez les noyaux d'attention linéaire (GDN) via Triton (--gdn-prefill-backend triton) pour éviter toute compilation à la volée (JIT) avec ninja/nvcc, et acheminez l'échantillonnage via torch natif (VLLM_USE_FLASHINFER_SAMPLER=0).

Tout ce qui se trouve ci-dessous s'exécute sur la même session Serverless GPU (H100) + environnement AI v6.

Enregistrer le modèle affiné pour le serving​

L'entraîneur conserve le modèle uniquement en mémoire (save_strategy="no"), pensez donc à le persister avant de redémarrer Python :

  1. save_pretrained les poids de texte fine-tuned (nommés model.*).
  2. Enregistrez le processor d'origine du modèle (preprocessor_config.json, etc.). La config composite déclare un composant de vision, donc vLLM l'exige. Réenregistrez ensuite le nouveau tokenizer par-dessus pour que le Template de chat personnalisé prime.
  3. Remplacer la configuration de texte plat par le Qwen3_5Configcomposite , avec architectures pin à Qwen3_5ForCausalLM pour que vLLM charge le modèle de texte (et non le modèle VL complet).

Exécutez ceci tant que sft_trainer, tokenizer et model_name sont toujours dans la portée.

Python
import os, tempfile
from transformers import AutoConfig, AutoProcessor

# Local-disk working dir. Use a fixed path (not a random tmpdir) so it survives %restart_python below;
# ARTIFACTS_PATH is a relative basename because the vLLM entrypoint's --model must match it both here
# and inside the packaged model's artifacts/ dir at serving time.
WORKDIR = os.path.join(tempfile.gettempdir(), "sft_serve")
ARTIFACTS_PATH = "qwen3_sft"
SAVE_DIR = os.path.join(WORKDIR, ARTIFACTS_PATH)
os.makedirs(SAVE_DIR, exist_ok=True)

# 1. Fine-tuned text backbone.
sft_trainer.model.save_pretrained(SAVE_DIR)

# 2. Original processor files, then the new tokenizer (with the custom chat template) on top.
try:
AutoProcessor.from_pretrained(model_name).save_pretrained(SAVE_DIR)
except Exception as e:
print("processor save skipped:", e)
tokenizer.save_pretrained(SAVE_DIR)

# 3. Composite Qwen3_5Config (text_config + vision_config), arch pinned to the text model.
cfg = AutoConfig.from_pretrained(model_name)
cfg.architectures = ["Qwen3_5ForCausalLM"]
cfg.save_pretrained(SAVE_DIR)

print("saved ->", SAVE_DIR)
print("config architectures:", cfg.architectures, "| model_type:", cfg.model_type,
"| has vision_config:", hasattr(cfg, "vision_config"))

Redémarrer Python pour libérer le H100​

Rien ne doit être installé : AI env v6 intègre déjà flashinfer et ses noyaux flashinfer-cubin précompilés, le préremplissage GDN (attention linéaire) s'exécute sur Triton, et l'échantillonnage s'exécute sur torch natif (tous deux définis dans les indicateurs de lancement ci-dessous), de sorte qu'aucun noyau ne subit de compilation JIT au Startup.

La seule chose nécessaire est de libérer la mémoire GPU retenue par le trainer afin que vLLM puisse l’utiliser. %restart_python effectue cette opération ; le point de contrôle enregistré sur le disque local persiste après le redémarrage (même nœud driver).

Python
%restart_python

Configuration de serving​

Le haut de la cellule suivante contient les valeurs que vous avez définies pour votre workspace : les noms des modèles/endpoints et le dimensionnement des endpoints. Le reste correspond au câblage interne que vous pouvez laisser tel quel. Au minimum, définissez UC_MODEL_NAME sur un chemin Unity Catalog dans lequel vous pouvez écrire avant d’exécuter les cellules d’enregistrement et de déploiement.

Python
from databricks.sdk.service.serving import ServingModelWorkloadType

# Re-read the widgets: %restart_python reset the Python process, but the widget
# values set at the top of the notebook persist and can be read again here.
UC_CATALOG = dbutils.widgets.get("uc_catalog")
UC_SCHEMA = dbutils.widgets.get("uc_schema")
UC_MODEL_NAME_BASE = dbutils.widgets.get("uc_model_name")
DEPLOY_ENDPOINT = dbutils.widgets.get("deploy_endpoint").lower() == "true"

# --- Model / endpoint names ---
UC_MODEL_NAME = f"{UC_CATALOG}.{UC_SCHEMA}.{UC_MODEL_NAME_BASE}" # Unity Catalog catalog.schema.model
ENDPOINT_NAME = f"{UC_MODEL_NAME_BASE}-endpoint" # serving endpoint name; unique per workspace
SERVED_MODEL_NAME = UC_MODEL_NAME_BASE # name vLLM exposes the model under

# --- Endpoint sizing (adjust if needed) ---
# --gdn-prefill-backend triton (see the entrypoint) JITs the GDN kernels for whatever GPU the pod
# lands on, so this is not pinned to Hopper. GPU_MEDIUM fits the 0.8B model; verify at deploy.
WORKLOAD_TYPE = ServingModelWorkloadType.GPU_MEDIUM
WORKLOAD_SIZE = "Small"
SCALE_TO_ZERO_ENABLED = True

# --- Internal wiring: leave as-is ---
import os, tempfile

# Local-disk working dir; must match the save cell (survives %restart_python; same driver node).
WORKDIR = os.path.join(tempfile.gettempdir(), "sft_serve")
ARTIFACTS_PATH = "qwen3_sft" # relative basename; entrypoint --model resolves to ./qwen3_sft
os.chdir(WORKDIR) # so `--model qwen3_sft` works locally and matches artifacts/ at serving

# Allowlisted ports for serverless GPU notebooks are 3000-3999. Model Serving requires 8080.
LOCAL_PORT = 3080
SERVING_PORT = 8080

# vLLM tuning.
DTYPE = "float16" # model is bf16-native; this matches what shipped
MAX_MODEL_LEN = 8192 # keep <= the model's max_position_embeddings (config.json)
GPU_MEMORY_UTILIZATION = 0.85

Testez le modèle affiné localement avec vLLM​

Démarrez un serveur vLLM OpenAI sur le point de contrôle enregistré et effectuez un test de fumée avant de consacrer environ 40 minutes au déploiement d'un endpoint. Un point de contrôle ou un indicateur incorrect échoue ici en quelques secondes à la place.

La commande entrypoint() est définie une seule fois et réutilisée à la fois pour le test local (LOCAL_PORT) et l'endpoint de mise en service (SERVING_PORT) ; la chaîne exacte est stockée dans les métadonnées du modèle et réexécutée à l'intérieur du conteneur de mise en service.

Python
def entrypoint(port: int) -> str:
args = [
"VLLM_USE_FLASHINFER_SAMPLER=0", # native torch sampling; avoids the flashinfer sampler JIT
"python", "-u", "-m", "vllm.entrypoints.openai.api_server",
"--model", ARTIFACTS_PATH,
"--served-model-name", SERVED_MODEL_NAME,
"--host", "0.0.0.0",
"--port", str(port),
"--dtype", DTYPE,
"--max-model-len", str(MAX_MODEL_LEN),
"--gpu-memory-utilization", str(GPU_MEMORY_UTILIZATION),
# Run GDN (linear-attention) prefill through Triton (its own bundled compiler) instead of
# flashinfer's JIT kernel, so no ninja/nvcc is needed here or in the serving pod.
"--gdn-prefill-backend", "triton",
# Text-only serving of a multimodal-config model: allow zero images/videos so vLLM never
# profiles or exercises the vision tower (absent from the text weights).
"--limit-mm-per-prompt", "'{\"image\": 0, \"video\": 0}'",
]
return " ".join(args)
Python
# Start the vLLM server in the background; logs stream to process.log.
import subprocess

log = open("process.log", "w")
subprocess.Popen(
["bash", "-lc", entrypoint(LOCAL_PORT)],
stdout=log,
stderr=subprocess.STDOUT,
text=True,
start_new_session=True,
)
Python
%sh
# Tail logs until vLLM is ready. If this hangs, vLLM startup probably hit an error (read process.log).
tail -f process.log | sed -u '/Application startup complete/q'
Python
# Smoke test (sync): the endpoint speaks the OpenAI chat schema at /invocations.
import requests

resp = requests.post(f"http://localhost:{LOCAL_PORT}/invocations", json={&quot;messages&quot;: [{&quot;role&quot;: &quot;user&quot;, &quot;content&quot;: &quot;Hello&quot;}]})
resp.json()["choices"][0]["message"]["content"]
Python
# Smoke test (streaming): vLLM streams completions as Server-Sent Events — one JSON chunk per
# `data: ` line, terminated by `data: [DONE]`.
import requests
import json

resp = requests.post(
f"http://localhost:{LOCAL_PORT}/invocations",
json={&quot;messages&quot;: [{&quot;role&quot;: &quot;user&quot;, &quot;content&quot;: &quot;Tell me a story that is about 300 words!&quot;}], &quot;stream&quot;: True},
stream=True,
)

for line in resp.iter_lines():
if not line:
continue
if line == b"data: [DONE]":
break
if line.startswith(b"data: "):
data = json.loads(line[6:])
delta = data["choices"][0].get("delta", {})
if "content" in delta:
print(delta["content"], end="", flush=True)
Python
%sh
# Stop the local server before logging/registering (the endpoint runs its own copy).
pkill -f vllm.entrypoints.openai.api_server

Étape 7 : Logs et enregistrement du modèle​

Enregistrez un ChatModel MLflow dont les métadonnées portent sur task = llm/v1/chat et la commande vLLM entrypoint. L'exécution du service exécute ce point d'entrée, et non python_model.predict, de sorte que le corps de la classe n'est qu'un espace réservé requis.

L’inscription avec env_pack="databricks_model_serving" génère les artefacts Serverless Optimized Deployment (SOD) (poids + environnement empaqueté) dont Custom LLM Serving a besoin. La journalisation doit être effectuée depuis ce Runtime GPU Serverless afin que les dépendances GPU appropriées soient empaquetées.

Python
import os
import mlflow
from mlflow.pyfunc.model import ChatModel, ChatCompletionResponse

# Required placeholder. Serving runs the entrypoint, not python_model.predict.
class LLMModel(ChatModel):
def predict(self, context, messages, params):
return ChatCompletionResponse.from_dict({"choices": []})

# You must log and register from a Serverless GPU runtime, otherwise the model is packaged
# with CPU deps and the GPU serving endpoint fails to start.
if not os.environ.get("DATABRICKS_ACCELERATOR"):
raise RuntimeError(
"This model MUST be logged+registered from a serverless GPU runtime, otherwise the correct dependencies will not be packaged for serving."
)

model_info = mlflow.pyfunc.log_model(
name=SERVED_MODEL_NAME,
python_model=LLMModel(),
artifacts={
&quot;model_dir&quot;: ARTIFACTS_PATH,
},
metadata={
&quot;task&quot;: &quot;llm/v1/chat&quot;,
&quot;entrypoint&quot;: entrypoint(SERVING_PORT),
},
# Pin whatever mlflow ships in AI env v6 (rather than a hardcoded version).
extra_pip_requirements=[f"mlflow=={mlflow.__version__}"],
)
model_info.model_uri
Python
import mlflow

# env_pack is required. Custom LLM Serving depends on Serverless Optimized Deployments (SOD).
# The endpoint will not work without it.
# https://docs.databricks.com/aws/en/machine-learning/model-serving/serverless-optimized-deployments
model_version = mlflow.register_model(model_info.model_uri, UC_MODEL_NAME, env_pack="databricks_model_serving")

Étape 8 : Créer l'endpoint de service​

Créer un endpoint Custom Foundation Model qui sert la version enregistrée. create_and_wait se bloque jusqu'à ce que l'endpoint soit prêt (jusqu'à 40 minutes pour le premier déploiement ; le conteneur download les artefacts SOD et amorce vLLM).

Déployer le endpoint (optionnel)​

Les étapes restantes déployant un endpoint de service géré. Ils nécessitent un service de modèle de fondation personnalisé basé sur un point d'entrée (déploiement optimisé Serverless). Si deploy_endpoint est false, le notebook s'arrête ici.

Python
if not DEPLOY_ENDPOINT:
dbutils.notebook.exit("deploy_endpoint=false: skipping managed serving endpoint deployment")
Python
from databricks.sdk import WorkspaceClient
from datetime import timedelta
from databricks.sdk.service.serving import EndpointCoreConfigInput, ServedEntityInput

served_entities = [
ServedEntityInput(
entity_name=UC_MODEL_NAME,
entity_version=str(model_version.version),
workload_type=WORKLOAD_TYPE,
workload_size=WORKLOAD_SIZE,
scale_to_zero_enabled=SCALE_TO_ZERO_ENABLED,
)
]

w = WorkspaceClient()

# Create the endpoint, or update it in place if an endpoint of this name already exists,
# so the notebook is safe to re-run. The first deploy can take up to ~40 minutes while the
# container downloads the Serverless Optimized Deployment artifacts and boots vLLM.
existing = {e.name for e in w.serving_endpoints.list()}
if ENDPOINT_NAME in existing:
print(f"Updating existing endpoint: {ENDPOINT_NAME}")
w.serving_endpoints.update_config_and_wait(
name=ENDPOINT_NAME, served_entities=served_entities, timeout=timedelta(minutes=40)
)
else:
print(f"Creating endpoint: {ENDPOINT_NAME}")
config = EndpointCoreConfigInput(name=ENDPOINT_NAME, served_entities=served_entities)
w.serving_endpoints.create_and_wait(
name=ENDPOINT_NAME, config=config, timeout=timedelta(minutes=40)
)

Étape 9 : interroger l'endpoint​

Une fois l’endpoint prêt, queryez-le de trois manières : le SDK Databricks, le client OpenAI et le client OpenAI avec streaming.

Python
# Query using the Databricks SDK.
from databricks.sdk import WorkspaceClient
from databricks.sdk.service.serving import ChatMessage, ChatMessageRole

w = WorkspaceClient()

resp = w.serving_endpoints.query(
name=ENDPOINT_NAME,
messages=[ChatMessage(role=ChatMessageRole.USER, content="Hi, what model are you?")],
)

print(resp.choices[0].message.content)
Python
# Query using the OpenAI client.
from openai import OpenAI

DATABRICKS_HOST = dbutils.notebook.entry_point.getDbutils().notebook().getContext().apiUrl().get()
DATABRICKS_TOKEN = dbutils.notebook.entry_point.getDbutils().notebook().getContext().apiToken().get()

client = OpenAI(
api_key=DATABRICKS_TOKEN,
base_url=f"{DATABRICKS_HOST}/serving-endpoints",
)

response = client.chat.completions.create(
model=ENDPOINT_NAME,
messages=[
{"role": "user", "content": "Hello"},
],
)
print(response.choices[0].message.content)
Python
# Query using the OpenAI client (streaming).
stream = client.chat.completions.create(
model=ENDPOINT_NAME,
messages=[
{"role": "user", "content": "Hello, tell me a 200 word story"},
],
stream=True,
)

for event in stream:
delta = event.choices[0].delta
print(delta.content, end="|")

Étapes suivantes​

Maintenant que vous avez ajusté, enregistré et servi votre modèle, vous pouvez :

Exemple de notebook​

Affinement supervisé (complet) et serving de Qwen3.5-0.8B