Views
No views yet
1from mlx_raclate.utils.utils import load
2from mlx_raclate.utils.token_classification import (
3 postprocess_token_classification_output,
4 viterbi_transition_biases_from_calibration,
5)
6
7# Load model and tokenizer
8model_path = "PITTI/pplx-embed-0.6b-nemotron"
9model, tokenizer = load(
10 model_path,
11 pipeline="token-classification"
12)
13
14# Prepare input texts
15texts = ['John works at Apple in California.', 'Microsoft was founded by Bill Gates.']
16
17# Tokenize
18max_length = getattr(model.config, "max_position_embeddings", 512)
19tokens = tokenizer._tokenizer(
20 texts,
21 return_tensors="mlx",
22 padding=True,
23 truncation=True,
24 max_length=max_length,
25 return_offsets_mapping=True,
26)
27offset_mapping = tokens.pop("offset_mapping")
28
29# Run inference
30outputs = model(
31 input_ids=tokens["input_ids"],
32 attention_mask=tokens["attention_mask"],
33 return_dict=True
34)
35
36# Get predictions
37logits = outputs["logits"]
38id2label = model.config.id2label
39transition_biases = viterbi_transition_biases_from_calibration(
40 getattr(model, "viterbi_calibration", None)
41)
42processed = postprocess_token_classification_output(
43 logits=logits,
44 probabilities=outputs["probabilities"],
45 id2label=id2label,
46 texts=texts,
47 offsets=offset_mapping.tolist(),
48 transition_biases=transition_biases,
49)
50
51# Process and print grouped spans
52for i, text in enumerate(texts):
53 print(f"Text: {text}")
54 print("Grouped spans:")
55 for span in processed["grouped_spans"][i]:
56 print(f" {span['entity_group']}: {span['word']!r} [{span['start']}, {span['end']}] score={span['score']:.3f}")
57 print()token-classification