small100-onnx / USAGE.md
casawolice's picture
Upload SMaLL-100 universal ONNX (int8 model, tokenizer.json, lang map, 4-platform examples)
a0d8498 verified
|
Raw
History Blame Contribute Delete
2.63 kB

Universal usage recipe

The same 4 steps work on every platform (transformers.js, Python, Android, iOS). Only the ONNX runtime and tokenizer binding differ; the algorithm is identical.

Artifacts

  • onnx/encoder_model.onnx — encoder (int8)
  • onnx/decoder_model_merged.onnx — decoder with KV cache, handles first + cached steps (int8)
  • tokenizer.json — HuggingFace fast tokenizer (validated; loads in tokenizers, transformers.js, DJL, swift-transformers)
  • lang_tokens.json — { "lang_to_id": {"en":128022,...}, "eos":2, "pad":1, "unk":3, "decoder_start":2 }

Algorithm

Given text and target language tgt (e.g. "en"):

  1. Encode + prepend target-language token (SMaLL-100 puts the target token on the source; the tokenizer already appends </s>=2):

    ids = [ lang_to_id[tgt] ] + tokenizer.encode(text).ids      # ends with 2 (</s>)
    attention_mask = [1] * len(ids)
    
  2. Run the encoder once:

    encoder_hidden_states = encoder(input_ids=ids, attention_mask)
    
  3. Greedy decode with the merged decoder (KV cache). Start from decoder_start_token_id = 2. The merged decoder takes a use_cache_branch flag: false on the first step (empty past), true afterwards.

    cur = 2 ; past = empty KV (3 layers × {decoder,encoder} × {key,value}, shape [1,16,0,64])
    use_cache = false ; out = []
    loop (max 128):
        logits, present = decoder(
            input_ids=[cur], encoder_hidden_states, encoder_attention_mask=attention_mask,
            past_key_values=past, use_cache_branch=use_cache)
        next = argmax(logits[0, -1])
        if next == 2 (</s>): break
        out.append(next)
        # roll KV: decoder.key/value always update; encoder.key/value only on step 1 then frozen
        past = present (decoder KV) ; keep encoder KV from step 1
        cur = next ; use_cache = true
    

    Model dims: 3 decoder layers, 16 heads, head_dim 64, d_model 1024. Decoder ONNX inputs: encoder_attention_mask, input_ids, encoder_hidden_states, past_key_values.{0..2}.{decoder,encoder}.{key,value}, use_cache_branch. Outputs: logits, present.{0..2}....

  4. Decode ids → text (skip special tokens):

    result = tokenizer.decode(out, skip_special_tokens=true)
    

optimum / transformers.js .generate() do steps 2–3 for you; on Android/iOS you run the loop yourself with onnxruntime (see examples/).

Language codes

lang_tokens.json maps 100 ISO codes → token ids (e.g. en, zh, ja, ko, fr, de, es, ru, ar, hi, th, vi, …). Pass the target code.