// Copyright 2025 International Digital Economy Academy
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
///|
struct Encoding {
encoder : Map[BytesView, Int]
decoder : FixedArray[Bytes]
special_encoder : Map[String, Int]
special_decoder : Map[Int, Bytes]
regex : @pcre2.Code
}
///|
pub fn Encoding::new(
mergeable_ranks~ : Map[Bytes, Int],
special_tokens~ : Map[String, Int],
pat_str~ : String,
) -> Encoding raise {
let encoder : Map[BytesView, Int] = {}
let decoder = FixedArray::make(mergeable_ranks.length(), b"")
for piece, rank in mergeable_ranks {
encoder[piece] = rank
decoder[rank] = piece
}
let special_encoder = {}
let special_decoder = {}
for token, rank in special_tokens {
special_encoder[token] = rank
special_decoder[rank] = @encoding/utf8.encode(token)
}
let regex = @pcre2.compile(pat_str)
regex.jit_compile(complete=true)
Encoding::{ encoder, decoder, special_encoder, special_decoder, regex }
}
///|
pub fn Encoding::special_tokens(self : Encoding, piece : String) -> Int? {
self.special_encoder.get(piece)
}
///|
fn arg_min(vec : Array[Int]) -> Int {
let mut value = @int.max_value
let mut index = -1
for i = 0; i < vec.length(); i = i + 1 {
if vec[i] < value {
value = vec[i]
index = i
}
}
index
}
///|
fn Encoding::lookup(
self : Encoding,
piece : Bytes,
tokens : Array[Int],
) -> Unit {
for i in 0.. index
None => @int.max_value
}
tokens.push(token)
}
}
///|
fn Encoding::rank(self : Encoding, start : Int, tokens : Array[Int]) -> Int {
let ranks = []
for i = start; i < tokens.length() - 1; i = i + 1 {
let l = tokens[i]
let r = tokens[i + 1]
let l = self.decoder[l]
let r = self.decoder[r]
let rank = match self.encoder.get(l + r) {
Some(merge) => merge
None => @int.max_value
}
ranks.push(rank)
}
arg_min(ranks)
}
///|
fn Encoding::merge(self : Encoding, start : Int, tokens : Array[Int]) -> Unit {
for index = self.rank(start, tokens)
index != -1
index = self.rank(start, tokens) {
let index = start + index
let l = tokens[index]
let r = tokens[index + 1]
let l = self.decoder[l]
let r = self.decoder[r]
let merge = l + r
let merge = self.encoder[merge]
tokens[index] = merge
tokens.remove(index + 1) |> ignore()
}
}
///|
pub fn Encoding::encode(
self : Encoding,
piece : StringView,
) -> Array[Int] raise {
let tokens = []
let matches = self.regex.matches(piece)
while matches.next() is Some(matched) {
let piece = @encoding/utf8.encode(matched[0])
match self.encoder.get(piece) {
Some(token) => tokens.push(token)
None => {
let start = tokens.length()
self.lookup(piece, tokens)
self.merge(start, tokens)
}
}
}
tokens
}
///|
suberror DecodingError {
InvalidToken(Int)
MalformedUtf8(Bytes)
} derive(Show)
///|
pub fn Encoding::decode(
self : Encoding,
tokens : ArrayView[Int],
) -> String raise DecodingError {
let buffer = @buffer.new()
for token in tokens {
if self.decoder.get(token) is Some(piece) {
buffer.write_bytes(piece)
continue
}
if self.special_decoder.get(token) is Some(piece) {
buffer.write_bytes(piece)
continue
}
raise DecodingError::InvalidToken(token)
}
let contents = buffer.contents()
@encoding/utf8.decode(contents) catch {
_ => raise DecodingError::MalformedUtf8(contents)
}
}