Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
7 changes: 6 additions & 1 deletion keras_hub/src/utils/transformers/convert_qwen3_5_moe.py
Original file line number Diff line number Diff line change
Expand Up @@ -36,6 +36,8 @@ def load_image_converter_config(preset, transformers_config):
"""Return kwargs for Qwen3_5ImageConverter, or None for text-only."""
if "vision_config" not in transformers_config:
return None
if transformers_config.get("language_model_only", False):
return None

vision_config = transformers_config["vision_config"]
preprocessor_config = load_json(preset, "preprocessor_config.json")
Expand Down Expand Up @@ -67,6 +69,8 @@ def load_video_converter_config(preset, transformers_config):
"""
if "vision_config" not in transformers_config:
return None
if transformers_config.get("language_model_only", False):
return None

vision_config = transformers_config["vision_config"]
video_config = load_json(preset, "video_preprocessor_config.json")
Expand Down Expand Up @@ -102,6 +106,7 @@ def convert_backbone_config(transformers_config):

top_level_vision_config = transformers_config.get("vision_config", None)
top_level_hidden_size = transformers_config.get("hidden_size", None)
language_model_only = transformers_config.get("language_model_only", False)

if "text_config" in transformers_config:
transformers_config = transformers_config["text_config"]
Expand All @@ -120,7 +125,7 @@ def convert_backbone_config(transformers_config):

vision_encoder = None
vision_config = top_level_vision_config
if vision_config is not None:
if vision_config is not None and not language_model_only:
text_hidden = transformers_config.get(
"hidden_size", top_level_hidden_size
)
Expand Down
65 changes: 58 additions & 7 deletions tools/checkpoint_conversion/convert_qwen3_5_moe_checkpoints.py
Original file line number Diff line number Diff line change
Expand Up @@ -25,6 +25,8 @@
from absl import flags
from keras import ops
from PIL import Image
from transformers import AutoConfig
from transformers import AutoModelForCausalLM
from transformers import AutoModelForImageTextToText
from transformers import AutoProcessor
from transformers import AutoTokenizer
Comment thread
laxmareddyp marked this conversation as resolved.
Expand All @@ -42,6 +44,7 @@
"qwen3_5_moe_35b_a3b_base": "Qwen/Qwen3.5-35B-A3B-Base",
"qwen3_5_moe_35b_a3b": "Qwen/Qwen3.5-35B-A3B",
"qwen3_6_moe_35b_a3b": "Qwen/Qwen3.6-35B-A3B",
"qwen_agent_world_35b_a3b": "Qwen/Qwen-AgentWorld-35B-A3B",
}

IMAGE_URL = "http://images.cocodataset.org/val2017/000000039769.jpg"
Expand Down Expand Up @@ -101,7 +104,37 @@ def _count_keras_params(backbone):


# ---------------------------------------------------------------
# 1. Precompute HF outputs (before freeing HF model)
# 1a. Precompute text-only HF outputs (for language_model_only models)
# ---------------------------------------------------------------
def precompute_text_only_outputs(hf_model, hf_tokenizer):
"""Precompute text-only HF outputs (no vision components)."""
results = {}

hf_ids = hf_tokenizer(TEXT_PROMPT, return_tensors="np")["input_ids"]
results["text_token_ids"] = hf_ids

with torch.no_grad():
hf_out = hf_model(
input_ids=torch.tensor(hf_ids, dtype=torch.long).to(device),
)
results["text_logits"] = hf_out.logits.detach().cpu().float().numpy()

if not FLAGS.skip_generation:
with torch.no_grad():
hf_gen = hf_model.generate(
input_ids=torch.tensor(hf_ids, dtype=torch.long).to(device),
max_new_tokens=32,
do_sample=False,
)
results["text_generated"] = hf_tokenizer.decode(
hf_gen[0], skip_special_tokens=True
)

return results


# ---------------------------------------------------------------
# 1b. Precompute HF outputs (multimodal, before freeing HF model)
# ---------------------------------------------------------------
def precompute_hf_outputs(hf_model, hf_tokenizer, hf_preset):
"""Precompute all HF outputs needed for validation.
Expand Down Expand Up @@ -565,6 +598,12 @@ def save_preset(keras_model, preset_name):
# ---------------------------------------------------------------
# Main
# ---------------------------------------------------------------
def _is_text_only(hf_preset):
"""Check if a HuggingFace preset is text-only (no vision weights)."""
config = AutoConfig.from_pretrained(hf_preset)
return getattr(config, "language_model_only", False)


def main(_):
preset = FLAGS.preset
if preset not in PRESET_MAP:
Expand All @@ -574,21 +613,33 @@ def main(_):
)

hf_preset = PRESET_MAP[preset]
text_only = _is_text_only(hf_preset)

# --- Phase 1: Load HF model and precompute all outputs ---
print("-> Loading HF model...")
hf_model = AutoModelForImageTextToText.from_pretrained(
hf_preset,
device_map="cpu",
torch_dtype=torch.float32,
)
if text_only:
print(" (text-only model detected, using AutoModelForCausalLM)")
hf_model = AutoModelForCausalLM.from_pretrained(
hf_preset,
device_map="cpu",
torch_dtype=torch.float32,
)
else:
hf_model = AutoModelForImageTextToText.from_pretrained(
hf_preset,
device_map="cpu",
torch_dtype=torch.float32,
)
hf_model.eval()
hf_tokenizer = AutoTokenizer.from_pretrained(hf_preset)
hf_params = sum(p.numel() for p in hf_model.parameters())
print(f" HF model loaded: {hf_params:,} params")

print("\n-> Precomputing all HF outputs...")
hf_results = precompute_hf_outputs(hf_model, hf_tokenizer, hf_preset)
if text_only:
hf_results = precompute_text_only_outputs(hf_model, hf_tokenizer)
else:
hf_results = precompute_hf_outputs(hf_model, hf_tokenizer, hf_preset)
hf_results["hf_param_count"] = hf_params
print(" HF outputs precomputed!")

Expand Down
Loading