Image-Text-to-Text
MLX
Safetensors
mage_vl
vision-language-model
video-understanding
mage-vl
conversational
custom_code
8-bit precision
Instructions to use mlx-community/Mage-VL-8bit with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- MLX
How to use mlx-community/Mage-VL-8bit with MLX:
# Make sure mlx-vlm is installed # pip install --upgrade mlx-vlm from mlx_vlm import load, generate from mlx_vlm.prompt_utils import apply_chat_template from mlx_vlm.utils import load_config # Load the model model, processor = load("mlx-community/Mage-VL-8bit") config = load_config("mlx-community/Mage-VL-8bit") # Prepare input image = ["http://images.cocodataset.org/val2017/000000039769.jpg"] prompt = "Describe this image." # Apply chat template formatted_prompt = apply_chat_template( processor, config, prompt, num_images=1 ) # Generate output output = generate(model, processor, formatted_prompt, image) print(output) - Notebooks
- Google Colab
- Kaggle
- Local Apps Settings
- LM Studio
File size: 5,102 Bytes
0aff9ba | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 | from dataclasses import dataclass
import torch
import torch.nn as nn
import torch.nn.functional as F
from mamba_ssm.models.mixer_seq_simple import create_block
from transformers import Qwen3Config
from transformers.models.qwen3 import Qwen3ForCausalLM
class PreNet(nn.Module):
def __init__(self, d_code, d_model):
super().__init__()
self.fc3 = nn.Linear(d_code, d_model)
def forward(self, x):
return F.leaky_relu(self.fc3(x))
class PostNet(nn.Module):
def __init__(self, d_model, n_class):
super().__init__()
self.fc3 = nn.Linear(d_model, n_class)
def forward(self, x):
return self.fc3(F.leaky_relu(x))
@dataclass
class SSMConfig:
d_model: int = 2560
n_ssm: int = 1
class VideoMamba(nn.Module):
def __init__(self, config):
super().__init__()
self.ssms = nn.ModuleList(
[create_block(config.d_model, d_intermediate=0, layer_idx=i) for i in range(config.n_ssm)]
)
self.norm_fn = nn.LayerNorm(config.d_model)
def forward(self, embeds, inference_params=None):
hidden_states = embeds
residual = None
for ssm in self.ssms:
hidden_states, residual = ssm(
hidden_states, residual, inference_params=inference_params
)
residual = hidden_states + residual if residual is not None else hidden_states
return self.norm_fn(residual.to(dtype=self.norm_fn.weight.dtype))
class Qwen3ForCausalLMCls(Qwen3ForCausalLM):
def forward(self, inputs_embeds=None, labels=None, attention_mask=None, **kwargs):
outputs = self.model(inputs_embeds=inputs_embeds, attention_mask=attention_mask)
logits = self.lm_head(outputs.last_hidden_state).float()
loss = None
if labels is not None:
shift_logits = logits[..., :-1, :].contiguous().view(-1, self.config.vocab_size)
shift_labels = labels[..., 1:].contiguous().view(-1).to(shift_logits.device)
loss = nn.CrossEntropyLoss(
weight=torch.tensor([0.15, 0.85], device=shift_logits.device)
)(shift_logits, shift_labels)
return {"loss": loss, "logits": logits}
class ClsNet(nn.Module):
def __init__(self, hidden_size=2560, num_layers=4):
super().__init__()
config = Qwen3Config(
vocab_size=2,
hidden_size=hidden_size,
num_hidden_layers=num_layers,
num_attention_heads=32,
num_key_value_heads=8,
intermediate_size=12288,
head_dim=128,
max_position_embeddings=8192,
rms_norm_eps=1e-6,
tie_word_embeddings=False,
attention_bias=False,
)
self.cls_model = Qwen3ForCausalLMCls(config)
def forward(self, x, labels=None, attention_mask=None):
return self.cls_model(inputs_embeds=x, labels=labels, attention_mask=attention_mask)
class StreamMindGate(nn.Module):
def __init__(self, hidden_size=2560):
super().__init__()
self.pre_net = PreNet(hidden_size, hidden_size)
self.mamba_model = VideoMamba(SSMConfig(d_model=hidden_size))
self.post_net = PostNet(hidden_size, hidden_size)
self.cls_net = ClsNet(hidden_size=hidden_size, num_layers=4)
def perception_tokens(self, vision_tokens):
"""Convert [B,T,P,D] visual patches to one EPFE token per time step."""
x = vision_tokens.mean(dim=2)
batch, time, dim = x.shape
x = self.pre_net(x.reshape(batch * time, dim)).reshape(batch, time, dim)
x = self.mamba_model(x)
x = self.post_net(x.reshape(batch * time, dim)).reshape(batch, time, dim)
return x
def forward(self, vision_tokens, response_positions=None):
"""Return [B,T,2] silent/speak logits for every EPFE time step."""
tokens = self.perception_tokens(vision_tokens)
batch, time, dim = tokens.shape
target_ids = torch.zeros(batch, time, dtype=torch.long, device=tokens.device)
if response_positions is not None:
target_ids[:, torch.as_tensor(response_positions, device=tokens.device) - 1] = 1
targets = self.cls_net.cls_model.model.embed_tokens(
target_ids.reshape(batch * time)
)
pair = torch.stack((tokens.reshape(batch * time, dim), targets), dim=1)
rotary = self.cls_net.cls_model.model.rotary_emb
saved_inv_freq = rotary.inv_freq
try:
# Match the training checkpoint, where the full model (including
# non-persistent Qwen3 RoPE buffers) was cast to BF16.
rotary.inv_freq = rotary.inv_freq.to(pair.dtype)
output = self.cls_net(
pair,
attention_mask=torch.ones(pair.shape[:2], device=pair.device),
)
finally:
rotary.inv_freq = saved_inv_freq
# Autoregressive shift: position 0 predicts the target token at
# position 1, matching StreamMind's logits[..., :-1, :] evaluation.
return output["logits"][:, 0].reshape(batch, time, 2)
|