patdev commited on
Commit
93e9593
·
verified ·
1 Parent(s): f9e2591

correctifs TRT-LLM : mappage MXFP4A16 et peripherique mRoPE

Browse files
patches_trtllm/README.md ADDED
@@ -0,0 +1,54 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Correctifs TensorRT-LLM 1.3.0rc22 (23/08/2026)
2
+
3
+ Deux correctifs Python — aucun CUDA nécessaire — qui débloquent TensorRT-LLM sur nos modèles.
4
+ Testés sur RTX 6000 Ada (sm89), TRT-LLM 1.3.0rc22.
5
+
6
+ ## Le raisonnement
7
+
8
+ TRT-LLM refusait nos MoE à cause du **format d'entrée**, pas d'une absence de noyau :
9
+
10
+ | checkpoint | ce qui se passe |
11
+ |---|---|
12
+ | INT4 groupwise (AWQ) | routé vers `WInt4AFP8FusedMoEMethod` (chemin **W4A8**) → `input_scale` absent → `'NoneType' object has no attribute 'device'` |
13
+ | FP8 officiel d'Ornith | `QuantMode 536` = `FP8_ROWWISE \| PER_TOKEN \| PER_CHANNEL` → aucune méthode MoE → `Unsupported quantization mode` |
14
+ | **MXFP4** | `get_mxfp4_quant_algo()` renvoie `W4A16_MXFP4` pour `sm < 100` → **`WFP4A16FusedMoEMethod`, qui existe** |
15
+
16
+ La voie MXFP4 était donc ouverte — mais deux verrous restaient.
17
+
18
+ ## `patch_mxfp4.py` — mappage compressed-tensors → W4A16_MXFP4
19
+
20
+ `update_quant_config_from_compressed_tensors()` (dans `models/quant_config_utils.py`) gère le
21
+ FP8 par canal, le FP8 par blocs et le NVFP4 (`strategy: "tensor_group"`, groupe 16), mais pas le
22
+ MXFP4 weight-only que produit le préréglage `MXFP4A16` de `llmcompressor` :
23
+ `strategy: "group"`, `group_size: 32`, `input_activations: null`.
24
+
25
+ Deux défauts corrigés :
26
+ 1. la fonction fait `inputs_quant_config["strategy"]` **sans vérifier que `input_activations`
27
+ n'est pas `None`** → `TypeError` avant même d'atteindre un test utile ;
28
+ 2. aucune branche ne mappe le MXFP4 vers `QuantAlgo.W4A16_MXFP4`.
29
+
30
+ ## `patch_mrope.py` — périphérique du cache mRoPE
31
+
32
+ `_prepare_qwen_vl_mrope_config()` (dans `_torch/models/modeling_qwen2vl.py`) construit `deltas`
33
+ par `torch.cat()` de tenseurs venant des paramètres multimodaux — donc du **CPU** — puis les écrit
34
+ dans `mrope_position_deltas_cache`, qui vit sur le **GPU** :
35
+
36
+ ```
37
+ RuntimeError: Expected all tensors to be on the same device,
38
+ but got source is on cpu, different from other tensors on cuda:0
39
+ ```
40
+
41
+ Le modèle charge entièrement et l'exécuteur démarre ; c'est la première génération qui tombe
42
+ dessus, ce qui rend le symptôme (`AssertionError: Sampling failed`) trompeur.
43
+
44
+ ## Utilisation
45
+
46
+ ```bash
47
+ python patch_mxfp4.py # idempotent, sauvegarde en .orig
48
+ python patch_mrope.py
49
+ python quant_mxfp4.py # produit un checkpoint MXFP4A16 depuis du BF16
50
+ python trt_mxfp4.py # sert le MoE MXFP4 en TRT-LLM
51
+ ```
52
+
53
+ `quant_mxfp4.py` laisse les portes de routage en pleine précision : quantifier un routeur dégrade
54
+ la sélection d'experts bien plus que les experts eux-mêmes.
patches_trtllm/patch_mrope.py ADDED
@@ -0,0 +1,25 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Correctif du bug de peripherique dans le mRoPE multimodal de TRT-LLM.
2
+ #
3
+ # _prepare_qwen_vl_mrope_config() construit `deltas` par torch.cat() de tenseurs
4
+ # venant des parametres multimodaux, qui vivent sur le CPU, puis les ecrit dans
5
+ # mrope_position_deltas_cache qui vit sur le GPU :
6
+ # RuntimeError: source is on cpu, different from other tensors on cuda:0
7
+ # Un .to(device) manquant. Le modele charge et l'executeur demarre : c'est la
8
+ # toute premiere generation qui tombe dessus.
9
+ import shutil, sys
10
+ F = "/opt/trtvenv/lib/python3.12/site-packages/tensorrt_llm/_torch/models/modeling_qwen2vl.py"
11
+ src = open(F, encoding="utf-8").read()
12
+ avant = " deltas = torch.cat(delta_tensors, dim=0)\n mrope_position_deltas_cache.index_copy_(0, seq_slots, deltas)"
13
+ apres = (" # les deltas viennent des parametres multimodaux, donc du CPU ;\n"
14
+ " # le cache vit sur le GPU -> aligner avant index_copy_\n"
15
+ " deltas = torch.cat(delta_tensors, dim=0).to(\n"
16
+ " device=mrope_position_deltas_cache.device,\n"
17
+ " dtype=mrope_position_deltas_cache.dtype)\n"
18
+ " mrope_position_deltas_cache.index_copy_(0, seq_slots, deltas)")
19
+ if apres.strip() in src:
20
+ print("PATCH_MROPE_DEJA"); sys.exit(0)
21
+ if avant not in src:
22
+ print("ANCRAGE_INTROUVABLE"); sys.exit(1)
23
+ shutil.copy(F, F + ".orig")
24
+ open(F, "w", encoding="utf-8").write(src.replace(avant, apres, 1))
25
+ print("PATCH_MROPE_APPLIQUE")
patches_trtllm/patch_mxfp4.py ADDED
@@ -0,0 +1,65 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Correctif : faire reconnaitre a TRT-LLM un checkpoint compressed-tensors
2
+ # MXFP4A16 (poids 4 bits flottants, groupe 32, activations non quantifiees).
3
+ #
4
+ # update_quant_config_from_compressed_tensors() gere trois cas -- FP8 par canal,
5
+ # FP8 par blocs, NVFP4 (strategy "tensor_group", groupe 16) -- mais pas le
6
+ # MXFP4 weight-only produit par le prereglage MXFP4A16 de llmcompressor
7
+ # (strategy "group", groupe 32, input_activations = None).
8
+ #
9
+ # Deux defauts a corriger, dans cet ordre :
10
+ # 1. la fonction fait inputs_quant_config["strategy"] sans verifier que
11
+ # input_activations n'est pas None -> TypeError avant tout test utile ;
12
+ # 2. aucune branche ne mappe le MXFP4 vers QuantAlgo.W4A16_MXFP4, alors que
13
+ # c'est precisement l'algo que get_mxfp4_quant_algo() choisit pour sm < 100
14
+ # et qu'il route vers WFP4A16FusedMoEMethod, qui est implementee.
15
+ import re, shutil, sys
16
+
17
+ F = "/opt/trtvenv/lib/python3.12/site-packages/tensorrt_llm/models/quant_config_utils.py"
18
+ src = open(F, encoding="utf-8").read()
19
+ if "W4A16_MXFP4" in src:
20
+ print("PATCH_DEJA_APPLIQUE"); sys.exit(0)
21
+ shutil.copy(F, F + ".orig")
22
+
23
+ # 1. ne pas dereferencer input_activations quand il vaut None
24
+ avant = ''' weights_quant_strategy = weights_quant_config["strategy"]
25
+ inputs_quant_strategy = inputs_quant_config["strategy"]'''
26
+ apres = ''' weights_quant_strategy = weights_quant_config["strategy"]
27
+ # MXFP4A16 laisse input_activations a None (poids seuls) : ne pas
28
+ # dereferencer avant d'avoir teste ce cas.
29
+ inputs_quant_strategy = (
30
+ inputs_quant_config["strategy"] if inputs_quant_config is not None else None
31
+ )'''
32
+ assert avant in src, "ancrage 1 introuvable"
33
+ src = src.replace(avant, apres)
34
+
35
+ # 2. brancher le MXFP4 weight-only avant la branche NVFP4
36
+ ancre = ''' elif (
37
+ weights_quant_config["num_bits"] == 4
38
+ and weights_quant_config.get("type") == "float"
39
+ and weights_quant_strategy == "tensor_group"
40
+ ):'''
41
+ nouveau = ''' elif (
42
+ weights_quant_config["num_bits"] == 4
43
+ and weights_quant_config.get("type") == "float"
44
+ and weights_quant_strategy == "group"
45
+ and inputs_quant_config is None
46
+ ):
47
+ # llm-compressor MXFP4A16 : poids FP4 avec echelles E8M0 par groupe de
48
+ # 32, activations laissees en 16 bits. Sur sm < 100, c'est exactement ce
49
+ # que get_mxfp4_quant_algo() demande, et cela route vers
50
+ # WFP4A16FusedMoEMethod.
51
+ group_size = weights_quant_config["group_size"]
52
+ if group_size != 32:
53
+ raise ValueError(f"Unsupported group_size: {group_size}. Supported: 32 for MXFP4.")
54
+ quant_config.quant_algo = QuantAlgo.W4A16_MXFP4
55
+ quant_config.group_size = group_size
56
+ elif (
57
+ weights_quant_config["num_bits"] == 4
58
+ and weights_quant_config.get("type") == "float"
59
+ and weights_quant_strategy == "tensor_group"
60
+ ):'''
61
+ assert ancre in src, "ancrage 2 introuvable"
62
+ src = src.replace(ancre, nouveau, 1)
63
+
64
+ open(F, "w", encoding="utf-8").write(src)
65
+ print("PATCH_MXFP4_APPLIQUE")
patches_trtllm/quant_mxfp4.py ADDED
@@ -0,0 +1,29 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Quantification MXFP4A16 -- la cle du deblocage TensorRT-LLM.
2
+ #
3
+ # Rappel du raisonnement : get_mxfp4_quant_algo() de TRT-LLM renvoie, pour
4
+ # sm < 100 (donc notre Ada sm89), QuantAlgo.W4A16_MXFP4, qui route vers
5
+ # WFP4A16FusedMoEMethod -- une methode MoE 4 bits a activations 16 bits qui
6
+ # EXISTE, contrairement au W4A16 groupwise (absent) et au FP8 rowwise (refuse).
7
+ #
8
+ # On quantifie l'A12B et non Ornith : la question posee est architecturale
9
+ # ("TRT-LLM sait-il charger un MoE Qwen3.5 en MXFP4 sur sm89 ?"), et Ornith
10
+ # demanderait de retelecharger 70 Go de BF16 pour la meme reponse.
11
+ #
12
+ # Les portes de routage restent en pleine precision : quantifier un routeur
13
+ # degrade la selection d'experts bien plus que les experts eux-memes.
14
+ import torch
15
+ from llmcompressor import oneshot
16
+ from llmcompressor.modifiers.quantization import QuantizationModifier
17
+ from transformers import AutoModelForCausalLM, AutoTokenizer
18
+
19
+ SRC = "/workspace/a12b"
20
+ DST = "/workspace/a12b-mxfp4"
21
+
22
+ model = AutoModelForCausalLM.from_pretrained(SRC, dtype=torch.bfloat16, device_map="cpu")
23
+ tok = AutoTokenizer.from_pretrained(SRC)
24
+ recette = QuantizationModifier(
25
+ targets="Linear", scheme="MXFP4A16",
26
+ ignore=["lm_head", "re:.*mlp\.gate$", "re:.*shared_expert_gate$"])
27
+ oneshot(model=model, recipe=recette, output_dir=DST)
28
+ tok.save_pretrained(DST)
29
+ print("QUANT_MXFP4_OK")
patches_trtllm/trt_dense.py ADDED
@@ -0,0 +1,26 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # TRT-LLM sur Qwen3.8-27B-FP8 (dense, block-scales -> chemin supporte).
2
+ #
3
+ # Deux corrections par rapport aux essais precedents :
4
+ # 1. garde __main__ : MPI re-importe le script dans chaque worker et
5
+ # respawnait a l'infini -> MPI_ABORT.
6
+ # 2. budget VRAM : le modele FP8 fait 27 Go sur une carte de 49 Go. Par
7
+ # defaut TRT-LLM reserve le reste pour le cache KV et depasse -> cudaMalloc
8
+ # out of memory. On borne explicitement la fraction libre et la longueur.
9
+ import time
10
+
11
+ def main():
12
+ from tensorrt_llm import LLM, SamplingParams
13
+ from tensorrt_llm.llmapi import KvCacheConfig
14
+ llm = LLM(model="/workspace/qwen38-fp8",
15
+ max_seq_len=4096, max_batch_size=1,
16
+ kv_cache_config=KvCacheConfig(free_gpu_memory_fraction=0.25))
17
+ sp = SamplingParams(max_tokens=256, temperature=0.0)
18
+ prompts = ["Ecris une fonction Python qui calcule la mediane d'une liste, avec docstring."]
19
+ llm.generate(prompts, sp)
20
+ t0 = time.time(); outs = llm.generate(prompts, sp); dt = time.time() - t0
21
+ n = len(outs[0].outputs[0].token_ids)
22
+ print(f"TRTLLM_DENSE {n} jetons en {dt:.2f}s = {n/dt:.1f} tok/s")
23
+ print("TEXTE:", outs[0].outputs[0].text[:400].replace("\n", " | "))
24
+
25
+ if __name__ == '__main__':
26
+ main()
patches_trtllm/trt_mxfp4.py ADDED
@@ -0,0 +1,32 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # LE test decisif : TRT-LLM sait-il servir un MoE Qwen3.5 en MXFP4 sur sm89 ?
2
+ #
3
+ # Attendu d'apres la lecture du code :
4
+ # load_hf_quant_config() voit quant_method == "mxfp4"
5
+ # -> get_mxfp4_quant_algo(), sm89 < 100 -> QuantAlgo.W4A16_MXFP4
6
+ # -> CutlassFusedMoE._get_quant_method() -> WFP4A16FusedMoEMethod
7
+ #
8
+ # Trois issues possibles, toutes informatives :
9
+ # - ca charge et genere -> le deblocage est reel, on peut viser Ornith
10
+ # - "Unsupported quantization mode" -> le mapping HF->QuantMode ne prend pas
11
+ # le format compressed-tensors (il faudrait ecrire l'adaptateur)
12
+ # - autre erreur -> a diagnostiquer, mais on aura depasse le mur
13
+ import time, json, sys
14
+
15
+ def main():
16
+ d = "/workspace/a12b-mxfp4"
17
+ print("quantization_config du checkpoint:",
18
+ json.dumps(json.load(open(d + "/config.json")).get("quantization_config", {}), indent=1)[:400])
19
+ from tensorrt_llm import LLM, SamplingParams
20
+ from tensorrt_llm.llmapi import KvCacheConfig
21
+ llm = LLM(model=d, max_seq_len=4096, max_batch_size=1,
22
+ kv_cache_config=KvCacheConfig(free_gpu_memory_fraction=0.35))
23
+ sp = SamplingParams(max_tokens=128, temperature=0.0)
24
+ p = ["Ecris une fonction Python qui calcule la mediane d'une liste."]
25
+ llm.generate(p, sp)
26
+ t0 = time.time(); outs = llm.generate(p, sp); dt = time.time() - t0
27
+ n = len(outs[0].outputs[0].token_ids)
28
+ print(f"TRTLLM_MXFP4 {n} jetons en {dt:.2f}s = {n/dt:.1f} tok/s")
29
+ print("TEXTE:", outs[0].outputs[0].text[:300].replace("\n", " | "))
30
+
31
+ if __name__ == '__main__':
32
+ main()