fix qwen3.5 sanitize (#928)

This commit is contained in:
Awni Hannun
2026-02-24 17:04:43 -08:00
committed by GitHub
parent 179da774b1
commit 834fac934c
2 changed files with 3 additions and 10 deletions
+2 -4
View File
@@ -5,7 +5,6 @@ from typing import Any, Dict, List, Optional, Union
import mlx.core as mx
import mlx.nn as nn
from mlx.utils import tree_flatten, tree_unflatten
from .base import (
BaseModelArgs,
@@ -364,11 +363,10 @@ class Model(nn.Module):
)
def sanitize(self, weights):
weights = tree_unflatten(list(weights.items()))
weights = dict(tree_flatten(weights))
sanitized = {}
for key, value in weights.items():
if key.startswith("vision_tower") or key.startswith("model.visual"):
continue
if key.startswith("model.visual"):
continue
if key.startswith("model.language_model"):
+1 -6
View File
@@ -2,8 +2,6 @@
from dataclasses import dataclass
from mlx.utils import tree_flatten, tree_unflatten
from .base import BaseModelArgs
from .qwen3_5 import Model as Qwen3_5Model
@@ -23,12 +21,9 @@ class ModelArgs(BaseModelArgs):
class Model(Qwen3_5Model):
def sanitize(self, weights):
weights = tree_unflatten(list(weights.items()))
weights = dict(tree_flatten(weights))
new_weights = {}
for key, value in weights.items():
if key.startswith("model.visual"):
if key.startswith("vision_tower") or key.startswith("model.visual"):
continue
if key.startswith("model.language_model"):
key = key.replace("model.language_model", "language_model.model")