Skip to content
Draft
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
23 changes: 21 additions & 2 deletions aiu_fms_testing_utils/scripts/inference.py
Original file line number Diff line number Diff line change
@@ -1,3 +1,3 @@
# Standard
import argparse
import datetime
Expand All @@ -5,6 +5,7 @@
import itertools
import json
import os
from packaging import version
from pathlib import Path
import random
import time
Expand All @@ -21,7 +22,7 @@
from fms.utils import generation
from fms.utils.generation import pad_input_ids

from transformers import AutoTokenizer
from transformers import AutoTokenizer, AutoConfig


# This example script validates the LLaMA implementation by running inference on a couple of prompts.
Expand Down Expand Up @@ -586,7 +587,25 @@
dprint(model)
dprint("=" * 60 + "\n")

tokenizer = AutoTokenizer.from_pretrained(args.tokenizer)
if hasattr(model.config, "model_type") and model.config.model_type in ["mistral", "mistral3", "pixtral"]:
# Check transformer version in the model to see if the regex fix can be avoided
# Ref: https://github.com/huggingface/transformers/blob/de306e8e14672dd8392b4bd344054a6a18de8613/src/transformers/tokenization_utils_tokenizers.py#L1205
# Above reference uses 4.57.2 version, but we get validation error from transformers if we do that check,
# so adjusted version based on testing
transformers_version = AutoConfig.from_pretrained(
args.tokenizer
).transformers_version

if transformers_version and version.parse(transformers_version) < version.parse(
"4.52.4"
):
tokenizer = AutoTokenizer.from_pretrained(
args.tokenizer, fix_mistral_regex=True
)
else:
tokenizer = AutoTokenizer.from_pretrained(args.tokenizer)
else:
tokenizer = AutoTokenizer.from_pretrained(args.tokenizer)
model.eval()
torch.set_grad_enabled(False)
loading_model_time = time.time() - loading_model_time
Expand Down
30 changes: 24 additions & 6 deletions aiu_fms_testing_utils/utils/paged.py
Original file line number Diff line number Diff line change
Expand Up @@ -134,6 +134,16 @@ def generate(
"_kvcache_num_blocks_hint", (_MAX_BATCH * _MAX_CONTEXT_LENGTH) // BLOCK_SIZE
)

# TODO: Implement a more hollistic vision handling
text_nlayers = None

if hasattr(model.config, "text_config"):
nheads = model.config.text_config.nheads
text_nlayers = model.config.text_config.nlayers
else:
nheads = model.config.nheads
text_nlayers = model.config.nlayers

if hasattr(model, "head"):
model_dtype = model.head.weight.dtype
elif hasattr(model, "shared"):
Expand All @@ -142,7 +152,6 @@ def generate(
else:
model_dtype = torch.float32

nheads = model.config.nheads
if hasattr(model.config, "kvheads"):
kvheads = model.config.kvheads
elif hasattr(model.config, "multiquery_attn"):
Expand All @@ -160,9 +169,18 @@ def generate(
raise ValueError("model must have a distributed_strategy")

kvheads = kvheads // tensor_parallel_size if kvheads > 1 else kvheads
head_size = getattr(
model.config, "head_dim", model.config.emb_dim // model.config.nheads
)

if hasattr(model.config, "text_config"):
head_size = getattr(
model.config.text_config,
"head_dim",
model.config.text_config.emb_dim // model.config.text_config.nheads,
)
else:
head_size = getattr(
model.config, "head_dim", model.config.emb_dim // model.config.nheads
)

if "fp8" in kwargs["attn_name"]:
from fms_mo.aiu_addons.fp8.fp8_utils import ScaledTensor

Expand Down Expand Up @@ -193,7 +211,7 @@ def generate(
already_scaled,
),
)
for _ in range(model.config.nlayers)
for _ in range(text_nlayers)
]
else:
kwargs["past_key_value_states"] = [
Expand All @@ -205,7 +223,7 @@ def generate(
NUM_BLOCKS, BLOCK_SIZE, kvheads, head_size, dtype=model_dtype
),
)
for _ in range(model.config.nlayers)
for _ in range(text_nlayers)
]
kwargs["block_table"] = None
block_numbers = [i for i in range(NUM_BLOCKS)]
Expand Down
Loading