This commit is contained in:
Xinyuan Tong
2026-04-29 21:13:51 +08:00
committed by GitHub
parent 4c1eefca4f
commit 1279ae0787
8 changed files with 885 additions and 5 deletions
@@ -0,0 +1,398 @@
---
title: Mistral Medium 3.5
metatags:
description: "Deploy Mistral Medium 3.5 with SGLang - 128B dense flagship merged model with hybrid reasoning, 256K context, vision input, and FP8 quantization."
---
import { MistralMedium35Deployment } from '/src/snippets/autoregressive/mistral-medium-3-5-deployment.jsx';
## 1. Model Introduction
**Mistral Medium 3.5** is Mistral AI's first flagship **merged model** — a single dense 128B checkpoint that handles instruction following, reasoning, and coding in one set of weights. It replaces Mistral Medium 3.1 and Magistral in Le Chat, and replaces Devstral 2 in the Vibe coding agent. Reasoning effort is configurable per request, so the same model can answer a quick chat reply or work through a deep agentic run. The vision encoder was trained from scratch to handle variable image sizes and aspect ratios.
**Key Features:**
- **Dense 128B parameters** — no MoE, no MLA, plain GQA (96 heads, 8 KV heads, head_dim=128)
- **256K context window** — YARN RoPE scaling on top of the original 4K base
- **Hybrid Reasoning**: Toggle between instant reply and deep reasoning per request via `reasoning_effort` (`"none"` or `"high"`)
- **Vision**: Accepts text + image input; from-scratch encoder that handles variable image sizes/aspect ratios
- **Function Calling**: Native tool calling and JSON output
- **FP8 Native**: Released with FP8 e4m3 static-tensor quantization built in
- **Multilingual**: 24 supported languages including English, French, German, Spanish, Portuguese, Italian, Japanese, Korean, Russian, Chinese, Arabic, Persian, Indonesian, Malay, Nepali, Polish, Romanian, Serbian, Swedish, Turkish, Ukrainian, Vietnamese, Hindi, and Bengali
- **License**: Modified MIT (open for commercial and non-commercial use except for companies with large revenue)
**Architecture:**
- Mistral 3 backbone with YARN RoPE for 256K context
- Dense (no MoE), 128B parameters
- Standard GQA attention (not MLA)
- Pixtral-style vision encoder (48 layers, patch_size=14, spatial_merge=2, image_size=1540) trained from scratch
- Multimodal input: text + image
**Models:**
- **[mistralai/Mistral-Medium-3.5-128B](https://huggingface.co/mistralai/Mistral-Medium-3.5-128B)** (FP8)
The HuggingFace repo ships both the mistral native layout (`params.json` + `consolidated-*.safetensors`) and the HF layout (`config.json` + `model-*.safetensors`). SGLang auto-detects the format — the HF layout is preferred when both are present.
---
## 2. SGLang Installation
SGLang offers multiple installation methods. You can choose the most suitable installation method based on your hardware platform and requirements.
Please refer to the [official SGLang installation guide](../../../docs/get-started/install) for installation instructions.
---
## 3. Model Deployment
### 3.1 Basic Configuration
**Interactive Command Generator**: Use the configuration selector below to generate a launch command for Mistral Medium 3.5.
<MistralMedium35Deployment />
### 3.2 Configuration Tips
- **Tensor Parallelism**: Mistral Medium 3.5 FP8 (~130 GB) requires `--tp 4` on Hopper (H100/H200) and `--tp 2` on Blackwell (B200/B300).
- **Reasoning effort**: Reasoning depth is configurable per request via `reasoning_effort` (`"none"`, `"high"`). No restart required — toggle per call.
- **Recommended temperature**: `0.7` when `reasoning_effort="high"`. Anywhere from `0.0` to `0.7` when `reasoning_effort="none"`, depending on the task — lower for to-the-point answers, higher for creative output.
- **Context length vs memory**: The model has a 256K context window. If you are memory-constrained, lower `--context-length` (e.g. `32768`) and increase once things are stable.
- **Tool calling**: Enable `--tool-call-parser mistral` to activate native function calling support.
- **Reasoning parser**: Enable `--reasoning-parser mistral` to separate `reasoning_content` from the main response content.
- **System prompt**: The model ships with a recommended system prompt in `chat_template.jinja` and `SYSTEM_PROMPT.txt`. If you do not pass a system message yourself, the chat template injects Mistral's default (model identity, current date, tool-use guidelines). For full fidelity with Mistral's reference setup, load `SYSTEM_PROMPT.txt` from the HF repo and substitute `{name}`, `{today}`, `{yesterday}` (see Section 4.6).
---
## 4. Model Invocation
### 4.1 Thinking Mode
Mistral Medium 3.5 is a hybrid reasoning model. By default it does not produce a reasoning trace — pass `reasoning_effort="high"` to switch on the deep-reasoning path. Mistral recommends `temperature=0.7` for reasoning mode.
```python Example
from openai import OpenAI
client = OpenAI(
base_url="http://localhost:30000/v1",
api_key="EMPTY",
)
response = client.chat.completions.create(
model="mistralai/Mistral-Medium-3.5-128B",
messages=[
{"role": "user", "content": "Solve step by step: what is 17 × 23 + 144 / 12?"},
],
temperature=0.7,
extra_body={"reasoning_effort": "high"},
)
print("Reasoning:", response.choices[0].message.reasoning_content)
print("Answer:", response.choices[0].message.content)
```
**Output:**
```text Output
Reasoning: I need to follow the order of operations (PEMDAS/BODMAS): multiplication and
division before addition, evaluated left to right.
17 × 23: I'll break it as 17 × (20 + 3) = 340 + 51 = 391.
144 / 12 = 12.
Finally, 391 + 12 = 403.
Answer: **17 × 23 + 144 / 12 = 403**
Step by step:
1. 17 × 23 = 391
2. 144 / 12 = 12
3. 391 + 12 = 403
```
### 4.2 Instruct Mode (Reasoning Off)
To skip the reasoning trace and get a fast direct response, set `reasoning_effort="none"`. For instruct mode, Mistral recommends temperature in the `0.0``0.7` range depending on how creative the task is:
```python Example
from openai import OpenAI
client = OpenAI(
base_url="http://localhost:30000/v1",
api_key="EMPTY",
)
response = client.chat.completions.create(
model="mistralai/Mistral-Medium-3.5-128B",
messages=[
{"role": "user", "content": "What is the capital of France?"},
],
temperature=0.1,
extra_body={"reasoning_effort": "none"},
)
print(response.choices[0].message.content)
```
**Output:**
```text Output
The capital of France is **Paris**. It is one of the most famous and visited cities in
the world, known for its rich history, art, culture, and landmarks like the Eiffel Tower,
Louvre Museum, and Notre-Dame Cathedral.
```
### 4.3 Streaming with Reasoning
```python Example
from openai import OpenAI
client = OpenAI(
base_url="http://localhost:30000/v1",
api_key="EMPTY",
)
stream = client.chat.completions.create(
model="mistralai/Mistral-Medium-3.5-128B",
messages=[
{"role": "user", "content": "Explain the difference between async and threading in Python."},
],
temperature=0.7,
extra_body={"reasoning_effort": "high"},
stream=True,
)
print("=== Reasoning ===")
for chunk in stream:
delta = chunk.choices[0].delta
if hasattr(delta, "reasoning_content") and delta.reasoning_content:
print(delta.reasoning_content, end="", flush=True)
elif delta.content:
print("\n=== Response ===")
print(delta.content, end="", flush=True)
print()
```
### 4.4 Tool Calling
Mistral Medium 3.5 supports native function calling. Enable with `--tool-call-parser mistral`:
```python Example
from openai import OpenAI
client = OpenAI(
base_url="http://localhost:30000/v1",
api_key="EMPTY",
)
tools = [
{
"type": "function",
"function": {
"name": "get_weather",
"description": "Get the current weather for a city",
"parameters": {
"type": "object",
"properties": {
"location": {"type": "string", "description": "City name"},
"unit": {"type": "string", "enum": ["celsius", "fahrenheit"]},
},
"required": ["location"],
},
},
}
]
response = client.chat.completions.create(
model="mistralai/Mistral-Medium-3.5-128B",
messages=[{"role": "user", "content": "What's the weather in Paris?"}],
tools=tools,
tool_choice="auto",
)
tool_calls = response.choices[0].message.tool_calls
for tc in tool_calls:
print(f"Tool: {tc.function.name}")
print(f"Args: {tc.function.arguments}")
```
**Output:**
```text Output
Tool: get_weather
Args: {"location": "Paris"}
```
### 4.5 Vision (Image Input)
Mistral Medium 3.5 accepts image inputs alongside text. The vision encoder was retrained from scratch to handle variable image sizes and aspect ratios:
```python Example
from openai import OpenAI
client = OpenAI(
base_url="http://localhost:30000/v1",
api_key="EMPTY",
)
response = client.chat.completions.create(
model="mistralai/Mistral-Medium-3.5-128B",
messages=[
{
"role": "user",
"content": [
{"type": "text", "text": "Describe what you see in this image."},
{
"type": "image_url",
"image_url": {"url": "https://raw.githubusercontent.com/sgl-project/sglang/main/assets/logo.png"},
},
],
}
],
temperature=0.7,
extra_body={"reasoning_effort": "none"},
)
print(response.choices[0].message.content)
```
**Output:**
```text Output
The image features a stylized representation of the acronym "SGL." The letters
are large, bold, and orange with a brown outline, giving them a three-dimensional
effect. To the left of the letters, there is a graphic that resembles a neuron
or a node with connections, also in a similar orange and brown color scheme. The
node has a code symbol (</>) inside a square, suggesting a connection to
programming or technology.
```
### 4.6 Loading the Reference System Prompt
Mistral ships a `SYSTEM_PROMPT.txt` alongside the weights. The reference setup loads it from the HF repo and substitutes `{name}`, `{today}`, and `{yesterday}` at runtime so the model knows its identity and the current date. SGLang's chat template will inject a default system prompt if you omit one, but for full parity with Mistral's reference, load it explicitly:
```python Example
from datetime import datetime, timedelta
from huggingface_hub import hf_hub_download
from openai import OpenAI
MODEL = "mistralai/Mistral-Medium-3.5-128B"
def load_system_prompt(repo_id: str, filename: str = "SYSTEM_PROMPT.txt") -> str:
path = hf_hub_download(repo_id=repo_id, filename=filename)
today = datetime.today().strftime("%Y-%m-%d")
yesterday = (datetime.today() - timedelta(days=1)).strftime("%Y-%m-%d")
name = repo_id.split("/")[-1]
with open(path) as f:
return f.read().format(name=name, today=today, yesterday=yesterday)
client = OpenAI(base_url="http://localhost:30000/v1", api_key="EMPTY")
response = client.chat.completions.create(
model=MODEL,
messages=[
{"role": "system", "content": load_system_prompt(MODEL)},
{"role": "user", "content": "Write me a sentence where every word starts with the next letter in the alphabet — start with 'a' and end with 'z'."},
],
temperature=0.1,
extra_body={"reasoning_effort": "none"},
)
print(response.choices[0].message.content)
```
---
## 5. Benchmarks
Validation runs on 4× H200 with `--tp 4`, served via the `/v1/chat/completions` endpoint.
### 5.1 Accuracy Benchmarks
#### GSM8K
```bash Command
python3 benchmark/gsm8k/bench_sglang.py --port 30000
```
**Results:**
```text Output
Accuracy: 0.945
Invalid: 0.000
Latency: 13.594 s
Output throughput: 1560.660 token/s
```
#### MMMU
```bash Command
python3 benchmark/mmmu/bench_sglang.py --port 30000
```
**Results:**
```text Output
Overall accuracy: 0.586
```
### 5.2 Speed Benchmarks
#### Latency (Low Concurrency)
```bash Command
python3 -m sglang.bench_serving \
--backend sglang \
--dataset-name random \
--num-prompts 10 \
--max-concurrency 1 \
--random-input-len 1024 \
--random-output-len 512 \
--port 30000
```
**Results:**
```text Output
============ Serving Benchmark Result ============
Backend: sglang
Successful requests: 10
Benchmark duration (s): 38.86
Total input tokens: 6101
Total generated tokens: 2684
Output token throughput (tok/s): 69.07
Mean E2E Latency (ms): 3883.80
Median TTFT (ms): 95.90
Median TPOT (ms): 14.19
==================================================
```
#### Throughput (High Concurrency)
```bash Command
python3 -m sglang.bench_serving \
--backend sglang \
--dataset-name random \
--num-prompts 1000 \
--max-concurrency 100 \
--random-input-len 1024 \
--random-output-len 512 \
--port 30000
```
**Results:**
```text Output
============ Serving Benchmark Result ============
Backend: sglang
Successful requests: 1000
Benchmark duration (s): 117.28
Total input tokens: 512842
Total generated tokens: 262023
Output token throughput (tok/s): 2234.18
Total token throughput (tok/s): 6607.01
Mean E2E Latency (ms): 11303.79
Median TTFT (ms): 152.95
Median TPOT (ms): 42.53
==================================================
```
+1
View File
@@ -1057,6 +1057,7 @@
"pages": [
"cookbook/autoregressive/Mistral/Ministral-3",
"cookbook/autoregressive/Mistral/Mistral-Small-4",
"cookbook/autoregressive/Mistral/Mistral-Medium-3.5",
"cookbook/autoregressive/Mistral/Devstral-2"
]
},
@@ -0,0 +1,340 @@
export const MistralMedium35Deployment = () => {
const modelId = 'mistralai/Mistral-Medium-3.5-128B';
const options = {
hardware: {
name: 'hardware',
title: 'Hardware Platform',
items: [
{ id: 'h100', label: 'H100', default: false },
{ id: 'h200', label: 'H200', default: true },
{ id: 'b200', label: 'B200', default: false },
{ id: 'b300', label: 'B300', default: false },
],
},
reasoning: {
name: 'reasoning',
title: 'Reasoning Parser',
items: [
{ id: 'disabled', label: 'Disabled', default: false },
{ id: 'enabled', label: 'Enabled', default: true }
],
commandRule: (value) => value === 'enabled' ? '--reasoning-parser mistral' : null
},
toolcall: {
name: 'toolcall',
title: 'Tool Call Parser',
items: [
{ id: 'disabled', label: 'Disabled', default: false },
{ id: 'enabled', label: 'Enabled', default: true }
],
commandRule: (value) => value === 'enabled' ? '--tool-call-parser mistral' : null
},
};
// 128B dense FP8 ≈ 130GB, plus KV cache headroom
const modelConfigs = {
h100: { tp: 4 },
h200: { tp: 4 },
b200: { tp: 2 },
b300: { tp: 2 },
};
const generateCommand = (values) => {
const { hardware } = values;
const hwConfig = modelConfigs[hardware];
if (!hwConfig) return `# Error: Unknown hardware combination`;
const { tp } = hwConfig;
let cmd = `sglang serve --model-path ${modelId}`;
cmd += ` \\\n --tp ${tp}`;
Object.entries(options).forEach(([key, option]) => {
if (key === 'hardware') return;
if (option.commandRule) {
const rule = option.commandRule(values[key]);
if (rule) cmd += ` \\\n ${rule}`;
}
});
return cmd;
};
const getInitialState = () => {
const initialState = {};
Object.entries(options).forEach(([key, option]) => {
if (option.type === 'checkbox') {
initialState[key] = (option.items || [])
.filter((item) => item.default)
.map((item) => item.id);
return;
}
if (option.type === 'text') {
initialState[key] = option.default || '';
return;
}
let items = option.items || [];
if (option.getDynamicItems) {
const defaultValues = {};
Object.entries(options).forEach(([innerKey, innerOption]) => {
if (innerOption.type === 'checkbox') {
defaultValues[innerKey] = (innerOption.items || [])
.filter((item) => item.default)
.map((item) => item.id);
} else if (innerOption.type === 'text') {
defaultValues[innerKey] = innerOption.default || '';
} else if (innerOption.items && innerOption.items.length > 0) {
const defaultItem = innerOption.items.find((item) => item.default);
defaultValues[innerKey] = defaultItem ? defaultItem.id : innerOption.items[0].id;
}
});
items = option.getDynamicItems(defaultValues);
}
const defaultItem = items && items.find((item) => item.default);
initialState[key] = defaultItem ? defaultItem.id : items && items[0] ? items[0].id : '';
});
return initialState;
};
const [values, setValues] = useState(getInitialState);
const [isDark, setIsDark] = useState(false);
useEffect(() => {
const checkDarkMode = () => {
const html = document.documentElement;
const isDarkMode =
html.classList.contains('dark') ||
html.getAttribute('data-theme') === 'dark' ||
html.style.colorScheme === 'dark';
setIsDark(isDarkMode);
};
checkDarkMode();
const observer = new MutationObserver(checkDarkMode);
observer.observe(document.documentElement, {
attributes: true,
attributeFilter: ['class', 'data-theme', 'style'],
});
return () => observer.disconnect();
}, []);
const handleRadioChange = (optionName, value) => {
setValues((prev) => ({ ...prev, [optionName]: value }));
};
const handleCheckboxChange = (optionName, itemId, isChecked) => {
setValues((prev) => {
const currentValues = prev[optionName] || [];
if (isChecked) {
return { ...prev, [optionName]: [...currentValues, itemId] };
}
return {
...prev,
[optionName]: currentValues.filter((id) => id !== itemId),
};
});
};
const handleTextChange = (optionName, value) => {
setValues((prev) => ({ ...prev, [optionName]: value }));
};
const command = generateCommand(values);
const containerStyle = {
maxWidth: '900px',
margin: '0 auto',
display: 'flex',
flexDirection: 'column',
gap: '4px',
};
const cardStyle = {
padding: '8px 12px',
border: `1px solid ${isDark ? '#374151' : '#e5e7eb'}`,
borderLeft: `3px solid ${isDark ? '#E85D4D' : '#D45D44'}`,
borderRadius: '4px',
display: 'flex',
alignItems: 'center',
gap: '12px',
background: isDark ? '#1f2937' : '#fff',
};
const titleStyle = {
fontSize: '13px',
fontWeight: '600',
minWidth: '140px',
flexShrink: 0,
color: isDark ? '#e5e7eb' : 'inherit',
};
const itemsStyle = {
display: 'flex',
rowGap: '2px',
columnGap: '6px',
flexWrap: 'wrap',
alignItems: 'center',
flex: 1,
};
const labelBaseStyle = {
padding: '4px 10px',
border: `1px solid ${isDark ? '#9ca3af' : '#d1d5db'}`,
borderRadius: '3px',
cursor: 'pointer',
display: 'inline-flex',
flexDirection: 'column',
alignItems: 'center',
justifyContent: 'center',
fontWeight: '500',
fontSize: '13px',
transition: 'all 0.2s',
userSelect: 'none',
minWidth: '45px',
textAlign: 'center',
flex: 1,
background: isDark ? '#374151' : '#fff',
color: isDark ? '#e5e7eb' : 'inherit',
};
const checkedStyle = {
background: '#D45D44',
color: 'white',
borderColor: '#D45D44',
};
const disabledStyle = {
cursor: 'not-allowed',
opacity: 0.5,
};
const subtitleStyle = {
display: 'block',
fontSize: '9px',
marginTop: '1px',
lineHeight: '1.1',
opacity: 0.7,
};
const textInputStyle = {
flex: 1,
padding: '8px 10px',
borderRadius: '4px',
border: `1px solid ${isDark ? '#4b5563' : '#d1d5db'}`,
background: isDark ? '#111827' : '#fff',
color: isDark ? '#e5e7eb' : '#111827',
fontSize: '13px',
};
const commandDisplayStyle = {
flex: 1,
padding: '12px 16px',
background: isDark ? '#111827' : '#f5f5f5',
borderRadius: '6px',
fontFamily: "'Menlo', 'Monaco', 'Courier New', monospace",
fontSize: '12px',
lineHeight: '1.5',
color: isDark ? '#e5e7eb' : '#374151',
whiteSpace: 'pre-wrap',
overflowX: 'auto',
margin: 0,
border: `1px solid ${isDark ? '#374151' : '#e5e7eb'}`,
};
return (
<div style={containerStyle} className="not-prose">
{Object.entries(options).map(([key, option]) => {
if (option.condition && !option.condition(values)) {
return null;
}
const items = option.getDynamicItems ? option.getDynamicItems(values) : option.items || [];
return (
<div key={key} style={cardStyle}>
<div style={titleStyle}>{option.title}</div>
<div style={itemsStyle}>
{option.type === 'text' ? (
<input
type="text"
value={values[option.name] || ''}
placeholder={option.placeholder || ''}
onChange={(event) => handleTextChange(option.name, event.target.value)}
style={textInputStyle}
/>
) : option.type === 'checkbox' ? (
(option.items || []).map((item) => {
const isChecked = (values[option.name] || []).includes(item.id);
const isDisabled =
item.required ||
(typeof item.disabledWhen === 'function' && item.disabledWhen(values));
return (
<label
key={item.id}
title={item.disabledReason || ''}
style={{
...labelBaseStyle,
...(isChecked ? checkedStyle : {}),
...(isDisabled ? disabledStyle : {}),
}}
>
<input
type="checkbox"
checked={isChecked}
disabled={isDisabled}
onChange={(event) =>
handleCheckboxChange(option.name, item.id, event.target.checked)
}
style={{ display: 'none' }}
/>
{item.label}
{item.subtitle && (
<small
style={{
...subtitleStyle,
color: isChecked ? 'rgba(255,255,255,0.85)' : 'inherit',
}}
>
{item.subtitle}
</small>
)}
</label>
);
})
) : (
items.map((item) => {
const isChecked = values[option.name] === item.id;
const isDisabled = Boolean(item.disabled);
return (
<label
key={item.id}
title={item.disabledReason || ''}
style={{
...labelBaseStyle,
...(isChecked ? checkedStyle : {}),
...(isDisabled ? disabledStyle : {}),
}}
>
<input
type="radio"
name={option.name}
value={item.id}
checked={isChecked}
disabled={isDisabled}
onChange={() => !isDisabled && handleRadioChange(option.name, item.id)}
style={{ display: 'none' }}
/>
{item.label}
{item.subtitle && (
<small
style={{
...subtitleStyle,
color: isChecked ? 'rgba(255,255,255,0.85)' : 'inherit',
}}
>
{item.subtitle}
</small>
)}
</label>
);
})
)}
</div>
</div>
);
})}
<div style={cardStyle}>
<div style={titleStyle}>Run this Command:</div>
<pre style={commandDisplayStyle}>{command}</pre>
</div>
</div>
);
};
+4 -1
View File
@@ -499,7 +499,10 @@ class ModelConfig:
or "LongcatFlashForCausalLMNextN" in self.hf_config.architectures
or "DotsVLMForCausalLM" in self.hf_config.architectures
or "MistralLarge3ForCausalLM" in self.hf_config.architectures
or "PixtralForConditionalGeneration" in self.hf_config.architectures
or (
"PixtralForConditionalGeneration" in self.hf_config.architectures
and getattr(self.hf_text_config, "kv_lora_rank", None) is not None
)
or "MistralLarge3ForCausalLMEagle" in self.hf_config.architectures
or "KimiK25ForConditionalGeneration" in self.hf_config.architectures
):
+102
View File
@@ -13,19 +13,81 @@
# ==============================================================================
"""Inference-only Mistral model."""
import logging
from collections.abc import Iterable
from typing import List
import regex as re
import torch
from transformers.models.mistral3.modeling_mistral3 import Mistral3MultiModalProjector
from sglang.srt.managers.schedule_batch import MultimodalDataItem
from sglang.srt.models.llama import LlamaForCausalLM
logger = logging.getLogger(__name__)
class MistralForCausalLM(LlamaForCausalLM):
pass
class MistralForCausalLMMistralFormat(MistralForCausalLM):
"""Mistral GQA model loaded from mistral native format (params.json).
Handles weight name remapping from mistral native format to HF/Llama
format. This is the GQA counterpart to MistralLarge3ForCausalLM which
handles MLA models in mistral native format.
"""
# fmt: off
remapping = {
r"layers\.(\d+)\.attention_norm\.weight": r"model.layers.\1.input_layernorm.weight",
r"layers\.(\d+)\.attention\.wq\.(\w+)": r"model.layers.\1.self_attn.q_proj.\2",
r"layers\.(\d+)\.attention\.wk\.(\w+)": r"model.layers.\1.self_attn.k_proj.\2",
r"layers\.(\d+)\.attention\.wv\.(\w+)": r"model.layers.\1.self_attn.v_proj.\2",
r"layers\.(\d+)\.attention\.wo\.(\w+)": r"model.layers.\1.self_attn.o_proj.\2",
r"layers\.(\d+)\.ffn_norm\.weight": r"model.layers.\1.post_attention_layernorm.weight",
r"layers\.(\d+)\.feed_forward\.w1\.(\w+)": r"model.layers.\1.mlp.gate_proj.\2",
r"layers\.(\d+)\.feed_forward\.w2\.(\w+)": r"model.layers.\1.mlp.down_proj.\2",
r"layers\.(\d+)\.feed_forward\.w3\.(\w+)": r"model.layers.\1.mlp.up_proj.\2",
r"norm\.weight": "model.norm.weight",
r"tok_embeddings\.weight": "model.embed_tokens.weight",
r"output\.weight": "lm_head.weight",
}
# fmt: on
def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]):
return super().load_weights(self._remap_mistral_to_llama(weights))
def _remap_mistral_to_llama(
self, weights: Iterable[tuple[str, torch.Tensor]]
) -> Iterable[tuple[str, torch.Tensor]]:
"""Remap Mistral native format weight names to HF/Llama format."""
for name, loaded_weight in weights:
# Pass through weights already in HF/Llama layout so this loader
# tolerates mixed-format checkpoints (e.g. native body + HF-style
# multi_modal_projector weights spliced in by a parent class).
if name.startswith("model.") or name.startswith("lm_head."):
yield name, loaded_weight
continue
for k, v in self.remapping.items():
match = re.fullmatch(k, name)
if match:
name = match.expand(v)
break
else:
logger.warning(f"Unrecognized weight: {name}. Skipping.")
continue
if name.endswith(".qscale_act"):
name = re.sub(r"\.qscale_act$", ".input_scale", name)
elif name.endswith(".qscale_weight"):
name = re.sub(r"\.qscale_weight$", ".weight_scale", name)
yield name, loaded_weight
class Mistral3ForConditionalGeneration:
MULTIMODAL_PROJECTOR_TYPE = Mistral3MultiModalProjector
@@ -89,5 +151,45 @@ class Mistral3ForConditionalGeneration:
def __call__(self, *args, **kwargs):
return self.inner(*args, **kwargs)
def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]):
"""Normalize transformers v5 Mistral3 weight names for
LlavaForConditionalGeneration.load_weights.
v5 checkpoints lay out Mistral3 weights as:
model.language_model.{embed_tokens,layers.*,norm}.*
model.vision_tower.*
model.multi_modal_projector.*
lm_head.*
The Llava loader routes by top-level `language_model.` /
`vision_tower.` prefixes, stripping one segment before forwarding to
the sub-module. The sub-module's own `load_weights` expects the
standard HF layout: `model.layers.*`, `model.embed_tokens.weight`,
`lm_head.weight` for Llama, and `vision_tower` internals at their
top level. So we rewrite:
model.language_model.X -> language_model.model.X
model.vision_tower.X -> vision_tower.X
model.multi_modal_projector.X -> multi_modal_projector.X
lm_head.X -> language_model.lm_head.X
"""
def normalize(ws):
for name, w in ws:
if name.startswith("model.language_model."):
rest = name[len("model.language_model.") :]
name = "language_model.model." + rest
elif name.startswith("model.vision_tower."):
name = "vision_tower." + name[len("model.vision_tower.") :]
elif name.startswith("model.multi_modal_projector."):
name = (
"multi_modal_projector."
+ name[len("model.multi_modal_projector.") :]
)
elif name.startswith("lm_head."):
name = "language_model." + name
yield name, w
return self.inner.load_weights(normalize(weights))
EntryClass = [MistralForCausalLM, Mistral3ForConditionalGeneration]
+16 -4
View File
@@ -45,6 +45,7 @@ from sglang.srt.managers.mm_utils import (
)
from sglang.srt.managers.schedule_batch import MultimodalDataItem, MultimodalInputs
from sglang.srt.model_loader.weight_utils import default_weight_loader
from sglang.srt.models.mistral import MistralForCausalLMMistralFormat
from sglang.srt.models.mistral_large_3 import MistralLarge3ForCausalLM
USE_XFORMERS_OPS = False
@@ -94,10 +95,21 @@ class PixtralForConditionalGeneration(nn.Module):
self.vision_args = VisionEncoderArgs(**vision_args)
self.language_model = MistralLarge3ForCausalLM(
config=self.config.text_config,
quant_config=kwargs.get("quant_config"),
)
# Choose language model based on text architecture:
# MLA text configs use DeepSeek V3 backbone (model_type="deepseek_v3"),
# GQA text configs use the standard Llama-style Mistral backbone.
text_config = self.config.text_config
is_mla = getattr(text_config, "model_type", "") == "deepseek_v3"
if is_mla:
self.language_model = MistralLarge3ForCausalLM(
config=text_config,
quant_config=kwargs.get("quant_config"),
)
else:
self.language_model = MistralForCausalLMMistralFormat(
config=text_config,
quant_config=kwargs.get("quant_config"),
)
self.vision_encoder = VisionTransformer(self.vision_args)
@@ -227,6 +227,14 @@ def get_hf_text_config(config: PretrainedConfig):
if getattr(_converted, "dtype", None) is None and parent_dtype is not None:
_converted.dtype = parent_dtype
setattr(config, _attr, _converted)
elif _sub is not None and parent_dtype is not None:
# transformers v5 multimodal configs (e.g. Mistral3Config) carry
# `dtype` only on the top-level config, leaving the sub-configs at
# None. Without this, _get_and_verify_dtype falls back to float32
# and then "auto" downcasts to float16, which overflows the Pixtral
# vision tower on real images and produces NaN features.
if getattr(_sub, "dtype", None) is None:
_sub.dtype = parent_dtype
# Priority: thinker_config > llm_config > language_config > text_config
if hasattr(config, "thinker_config"):
@@ -73,6 +73,22 @@ def adapt_config_dict(
config_dict["architectures"] = ["MixtralForCausalLM"]
else:
config_dict["architectures"] = ["MistralForCausalLM"]
config_dict["model_type"] = "mistral"
# Mistral models use non-interleaved RoPE (is_neox_style=False),
# unlike Llama which defaults to True.
config_dict["rope_is_neox_style"] = False
# Remove None-valued MLA fields that would shadow defaults in
# model_config._derive_model_shapes (getattr returns None instead
# of the fallback when the attribute exists but is None).
for mla_key in (
"q_lora_rank",
"qk_rope_head_dim",
"qk_nope_head_dim",
"kv_lora_rank",
"v_head_dim",
):
if config_dict.get(mla_key) is None:
config_dict.pop(mla_key, None)
if bool(config_dict.get("yarn")):
config_dict = _remap_mistral_yarn_args(config_dict)