bart-large-cnn-mlx

English | 日本語

Model Summary

This is an unofficial MLX conversion of facebook/bart-large-cnn (a standard BART encoder-decoder fine-tuned for summarization on CNN/DailyMail). All credit for the original model goes to its authors (Meta/Facebook AI).

This cannot be loaded with mlx_lm

mlx_lm only supports causal (decoder-only) language models, not encoder-decoder architectures like BART. This model was reimplemented from scratch for MLX and requires the bundled bart_mlx.py, which includes an autoregressive decode loop with KV caching.

Limitations

  • Greedy decoding only: the reference model's default generation config uses beam search (num_beams=4) plus a no-repeat-3-gram constraint; this implementation only supports greedy decoding. Output text may differ from the PyTorch beam-search reference (teacher-forced logits match exactly -- see Accuracy below).
  • No padding, single document only: batched processing of multiple documents is not supported.
  • forced_bos_token_id is applied: the reference generation config forces token ID 0 right after the decoder-start token. Skipping this noticeably degrades output quality (greedy decoding is very sensitive to the first token choice), so it's applied by default here.

Usage

import mlx.core as mx
from mlx.utils import tree_unflatten
from transformers import AutoTokenizer

from bart_mlx import BartMLX

tokenizer = AutoTokenizer.from_pretrained("facebook/bart-large-cnn")
model = BartMLX()
weights = mx.load("model.safetensors")
model.update(tree_unflatten(list(weights.items())))
mx.eval(model.parameters())

article = "..."  # the text to summarize
inputs = tokenizer(article, return_tensors="np")
input_ids = mx.array(inputs["input_ids"])

generated_ids = model.generate(input_ids, max_new_tokens=142)
summary = tokenizer.decode(generated_ids, skip_special_tokens=True)
print(summary)

Accuracy

Compared against the PyTorch fp32 reference (teacher-forced with the beam-search reference's own output sequence) on a real article about the Eiffel Tower:

Precision Teacher-forced logits cosine sim. Next-token argmax agreement
MLX fp32 1.0000001 100%
MLX fp16 (this release) 0.99999356 100%

Free-running greedy generation, compared against PyTorch's own greedy (num_beams=1) output:

  • PyTorch greedy: "The tower is 324 metres (1,063 ft) tall, about the same height as an 81-storey building. It is the tallest structure in Paris and the second tallest free-standing structure in France after the Millau Viaduct."
  • MLX greedy: generates the exact same text as above.

Specs

Item Value
Base model facebook/bart-large-cnn (BART-large, 406M params)
Precision float16
Framework MLX (from-scratch bart_mlx.py, autoregressive decoding with KV cache)

Notes

  • This is a community conversion, not an official release from Meta.
  • Security audit uses model-audit-lite (see SECURITY.md for details).

Security

Audited against its upstream with model-audit-lite: weight format, bundled code, and a machine-readable lineage (ML-BOM). Details, checksums and how to reproduce: SECURITY.md.


モデルの概要

facebook/bart-large-cnn(CNN/DailyMail で要約にファインチューニングされた、標準的なBART encoder-decoder)の MLX版です。元モデルの著作権はその作者(Meta/Facebook AI)に帰属します。

mlx_lmでは読み込めません

mlx_lmは因果言語モデル(decoder-only)専用で、BARTのようなencoder-decoderモデルには 対応していないため、MLXでの実装をゼロから書き起こして変換しています。同梱の bart_mlx.pyが必要です。KVキャッシュ付きの自己回帰デコードも実装済みです。

制約

  • グリーディデコードのみ: 元モデルの既定の生成設定はビームサーチ(num_beams=4)+ 3-gramの繰り返し禁止ですが、本実装はグリーディ(貪欲法)のみ対応です。 PyTorch版のビームサーチ結果とは文章として異なる場合があります(teacher forcingでの ロジット比較では完全に一致することを確認済み — 下記精度検証を参照)。
  • パディング無し・単一文書のみ: バッチで複数文書をまとめて処理する使い方は未対応です。
  • forced_bos_token_idの適用: 元モデルの生成設定では、デコーダー開始トークンの直後に 強制的にトークンID 0を出力する設定になっています。これを省略すると出力が著しく劣化する (グリーディデコードは最初のトークン選択に非常に敏感)ため、デフォルトで適用しています。

使い方

import mlx.core as mx
from mlx.utils import tree_unflatten
from transformers import AutoTokenizer

from bart_mlx import BartMLX

tokenizer = AutoTokenizer.from_pretrained("facebook/bart-large-cnn")
model = BartMLX()
weights = mx.load("model.safetensors")
model.update(tree_unflatten(list(weights.items())))
mx.eval(model.parameters())

article = "..."  # the text to summarize
inputs = tokenizer(article, return_tensors="np")
input_ids = mx.array(inputs["input_ids"])

generated_ids = model.generate(input_ids, max_new_tokens=142)
summary = tokenizer.decode(generated_ids, skip_special_tokens=True)
print(summary)

精度検証

PyTorch fp32リファレンス(ビームサーチ出力をteacher forcingで本実装に与えてロジットを比較) と、エッフェル塔に関する記事1件で比較:

精度 Teacher-forcedロジットのコサイン類似度 次トークンargmax一致率
MLX fp32 1.0000001 100%
MLX fp16(本リリース) 0.99999356 100%

グリーディ生成の実文章比較(PyTorch版もnum_beams=1でグリーディにした場合):

  • PyTorch greedy: "The tower is 324 metres (1,063 ft) tall, about the same height as an 81-storey building. It is the tallest structure in Paris and the second tallest free-standing structure in France after the Millau Viaduct."
  • MLX greedy: 上記と完全に同一の文章を生成。

Specs

Item Value
ベースモデル facebook/bart-large-cnn(BART-large、406M params)
精度 float16
フレームワーク MLX(ゼロから実装したbart_mlx.py、KVキャッシュ付き自己回帰デコード)

備考

  • 本変換は非公式のコミュニティ版です。
  • セキュリティー監査にはmodel-audit-liteを 使用しています(詳細はSECURITY.md)。

セキュリティー

model-audit-lite で変換元と突き合わせて監査済みです(重みの形式、同梱コード、機械可読な系譜=ML-BOM)。詳細・チェックサム・再現方法は SECURITY.md をご覧ください。

Downloads last month

-

Downloads are not tracked for this model. How to track
Safetensors
Model size
0.5B params
Tensor type
F16
·
MLX
Hardware compatibility
Log In to add your hardware

Quantized

Inference Providers NEW
This model isn't deployed by any Inference Provider. 🙋 Ask for provider support

Model tree for masahiroid/bart-large-cnn-mlx

Finetuned
(444)
this model