// Licensed to the .NET Foundation under one or more agreements.
// The .NET Foundation licenses this file to you under the MIT license.
// See the LICENSE file in the project root for more information.
using System;
using System.Buffers;
using System.Collections.Generic;
namespace Microsoft.ML.Tokenizers
{
///
/// Provides an abstraction for tokenizers, enabling the encoding of text into tokens and the decoding of token IDs back into text.
///
public abstract class Tokenizer
{
///
/// Initializes a new instance of the class.
///
protected Tokenizer() { }
///
/// Gets the PreTokenizer used by the Tokenizer.
///
public virtual PreTokenizer? PreTokenizer => null;
///
/// Gets the Normalizer in use by the Tokenizer.
///
public virtual Normalizer? Normalizer => null;
///
/// Encodes input text to token Ids.
///
/// The text to encode.
/// The span of the text to encode which will be used if the is .
/// The settings used to encode the text.
/// The encoded results containing the list of encoded Ids.
///
/// Types derived from may override this implementation to provide a more efficient implementation.
/// By default, it uses .
///
protected virtual EncodeResults EncodeToIds(string? text, ReadOnlySpan textSpan, EncodeSettings settings)
{
EncodeResults results = EncodeToTokens(text, textSpan, settings);
var ids = new int[results.Tokens.Count];
for (int i = 0; i < ids.Length; i++)
{
ids[i] = results.Tokens[i].Id;
}
return new EncodeResults
{
Tokens = ids,
CharsConsumed = results.CharsConsumed,
NormalizedText = results.NormalizedText,
};
}
///
/// Encodes input text to token Ids.
///
/// The text to encode.
/// Indicate whether to consider pre-tokenization before tokenization.
/// Indicate whether to consider normalization before tokenization.
/// The list of encoded Ids.
public IReadOnlyList EncodeToIds(string text, bool considerPreTokenization = true, bool considerNormalization = true)
=> EncodeToIds(text, text.AsSpan(), new EncodeSettings { ConsiderPreTokenization = considerPreTokenization, ConsiderNormalization = considerNormalization }).Tokens;
///
/// Encodes input text to token Ids.
///
/// The text to encode.
/// Indicate whether to consider pre-tokenization before tokenization.
/// Indicate whether to consider normalization before tokenization.
/// The list of encoded Ids.
public IReadOnlyList EncodeToIds(ReadOnlySpan text, bool considerPreTokenization = true, bool considerNormalization = true)
=> EncodeToIds(null, text, new EncodeSettings { ConsiderPreTokenization = considerPreTokenization, ConsiderNormalization = considerNormalization }).Tokens;
///
/// Encodes input text to token Ids up to maximum number of tokens.
/// The text to encode.
///
/// The maximum number of tokens to encode.
/// If the tokenizer's normalization is enabled or is , this will be set to in its normalized form; otherwise, this value will be set to .
/// The characters count of the text that encompasses the maximum encoded tokens.
/// Indicate whether to consider pre-tokenization before tokenization.
/// Indicate whether to consider normalization before tokenization.
/// The list of encoded Ids.
public IReadOnlyList EncodeToIds(string text, int maxTokenCount, out string? normalizedText, out int charsConsumed, bool considerPreTokenization = true, bool considerNormalization = true)
{
EncodeResults result = EncodeToIds(text, text.AsSpan(),
new EncodeSettings
{
ConsiderPreTokenization = considerPreTokenization,
ConsiderNormalization = considerNormalization,
MaxTokenCount = maxTokenCount
});
normalizedText = result.NormalizedText;
charsConsumed = result.CharsConsumed;
return result.Tokens;
}
///
/// Encodes input text to token Ids up to maximum number of tokens.
///
/// The text to encode.
/// The maximum number of tokens to encode.
/// If the tokenizer's normalization is enabled or is , this will be set to in its normalized form; otherwise, this value will be set to .
/// The characters count of the text that encompasses the maximum encoded tokens.
/// Indicate whether to consider pre-tokenization before tokenization.
/// Indicate whether to consider normalization before tokenization.
/// The list of encoded Ids.
public IReadOnlyList EncodeToIds(ReadOnlySpan text, int maxTokenCount, out string? normalizedText, out int charsConsumed, bool considerPreTokenization = true, bool considerNormalization = true)
{
EncodeResults result = EncodeToIds(null, text,
new EncodeSettings
{
ConsiderPreTokenization = considerPreTokenization,
ConsiderNormalization = considerNormalization,
MaxTokenCount = maxTokenCount
});
normalizedText = result.NormalizedText;
charsConsumed = result.CharsConsumed;
return result.Tokens;
}
///
/// Encodes input text to a list of s.
///
/// The text to encode.
/// The span of the text to encode which will be used if the is .
/// The settings used to encode the text.
protected abstract EncodeResults EncodeToTokens(string? text, ReadOnlySpan textSpan, EncodeSettings settings);
///
/// Encodes input text to a list of s.
///
/// The text to encode.
/// If the tokenizer's normalization is enabled or is , this will be set to in its normalized form; otherwise, this value will be set to .
/// Indicate whether to consider pre-tokenization before tokenization.
/// Indicate whether to consider normalization before tokenization.
/// The list of encoded s.
public IReadOnlyList EncodeToTokens(string text, out string? normalizedText, bool considerPreTokenization = true, bool considerNormalization = true)
{
EncodeResults result = EncodeToTokens(text, text.AsSpan(), new EncodeSettings { ConsiderPreTokenization = considerPreTokenization, ConsiderNormalization = considerNormalization });
normalizedText = result.NormalizedText;
return result.Tokens;
}
///
/// Encodes input text to a list of s.
///
/// The text to encode.
/// If the tokenizer's normalization is enabled or is , this will be set to in its normalized form; otherwise, this value will be set to .
/// Indicate whether to consider pre-tokenization before tokenization.
/// Indicate whether to consider normalization before tokenization.
/// The list of encoded s.
public IReadOnlyList EncodeToTokens(ReadOnlySpan text, out string? normalizedText, bool considerPreTokenization = true, bool considerNormalization = true)
{
EncodeResults result = EncodeToTokens(null, text, new EncodeSettings { ConsiderPreTokenization = considerPreTokenization, ConsiderNormalization = considerNormalization });
normalizedText = result.NormalizedText;
return result.Tokens;
}
///
/// Get the number of tokens that the input text will be encoded to.
///
/// The text to encode.
/// The span of the text to encode which will be used if the is .
/// The settings used to encode the text.
/// The number of token Ids that the input text will be encoded to.
///
/// Types derived from may override this implementation to provide a more efficient implementation.
/// By default, it uses .
///
protected virtual int CountTokens(string? text, ReadOnlySpan textSpan, EncodeSettings settings)
=> EncodeToTokens(text, textSpan, settings).Tokens.Count;
///
/// Get the number of tokens that the input text will be encoded to.
///
/// The text to encode.
/// Indicate whether to consider pre-tokenization before tokenization.
/// Indicate whether to consider normalization before tokenization.
/// The number of token Ids that the input text will be encoded to.
public int CountTokens(string text, bool considerPreTokenization = true, bool considerNormalization = true)
=> CountTokens(text, text.AsSpan(), new EncodeSettings { ConsiderPreTokenization = considerPreTokenization, ConsiderNormalization = considerNormalization });
///
/// Get the number of tokens that the input text will be encoded to.
///
/// The text to encode.
/// Indicate whether to consider pre-tokenization before tokenization.
/// Indicate whether to consider normalization before tokenization.
/// The number of token Ids that the input text will be encoded to.
public int CountTokens(ReadOnlySpan text, bool considerPreTokenization = true, bool considerNormalization = true)
=> CountTokens(null, text, new EncodeSettings { ConsiderPreTokenization = considerPreTokenization, ConsiderNormalization = considerNormalization });
///
/// Find the index of the maximum encoding capacity without surpassing the token limit.
///
/// The text to encode.
/// The span of the text to encode which will be used if the is .
/// The settings used to encode the text.
/// Indicate whether to find the index from the end of the text.
/// If the tokenizer's normalization is enabled or has is , this will be set to in its normalized form; otherwise, this value will be set to .
/// The token count can be generated which should be smaller than the maximum token count.
///
/// The index of the maximum encoding capacity within the processed text without surpassing the token limit.
/// If is , it represents the index immediately following the last character to be included. In cases where no tokens fit, the result will be 0; conversely,
/// if all tokens fit, the result will be length of the input text or the if the normalization is enabled.
/// If is , it represents the index of the first character to be included. In cases where no tokens fit, the result will be the text length; conversely,
/// if all tokens fit, the result will be zero.
///
///
/// Types derived from may override this implementation to provide a more efficient implementation.
/// By default, it uses .
///
protected virtual int GetIndexByTokenCount(string? text, ReadOnlySpan textSpan, EncodeSettings settings, bool fromEnd, out string? normalizedText, out int tokenCount)
{
int maxTokenCount = settings.MaxTokenCount;
if (fromEnd)
{
// If we're looking from the end, we need to process the whole input.
settings.MaxTokenCount = int.MaxValue;
}
EncodeResults tokens = EncodeToTokens(text, textSpan, settings);
normalizedText = tokens.NormalizedText;
tokenCount = Math.Min(maxTokenCount, tokens.Tokens.Count);
if (!fromEnd)
{
if (tokenCount > 0)
{
var token = tokens.Tokens[tokenCount - 1];
return token.Offset.End.Value;
}
return 0;
}
else
{
if (tokenCount > 0)
{
var token = tokens.Tokens[tokens.Tokens.Count - tokenCount];
return token.Offset.Start.Value;
}
return tokens.NormalizedText?.Length ?? textSpan.Length;
}
}
///
/// Find the index of the maximum encoding capacity without surpassing the token limit.
///
/// The text to encode.
/// The maximum number of tokens to encode.
/// If the tokenizer's normalization is enabled or is , this will be set to in its normalized form; otherwise, this value will be set to .
/// The token count can be generated which should be smaller than the maximum token count.
/// Indicate whether to consider pre-tokenization before tokenization.
/// Indicate whether to consider normalization before tokenization.
///
/// The index of the maximum encoding capacity within the processed text without surpassing the token limit.
/// It represents the index immediately following the last character to be included. In cases where no tokens fit, the result will be 0; conversely,
/// if all tokens fit, the result will be length of the input text or the if the normalization is enabled.
///
public int GetIndexByTokenCount(string text, int maxTokenCount, out string? normalizedText, out int tokenCount, bool considerPreTokenization = true, bool considerNormalization = true)
=> GetIndexByTokenCount(
text,
text.AsSpan(),
new EncodeSettings { ConsiderPreTokenization = considerPreTokenization, ConsiderNormalization = considerNormalization, MaxTokenCount = maxTokenCount },
fromEnd: false,
out normalizedText,
out tokenCount);
///
/// Find the index of the maximum encoding capacity without surpassing the token limit.
///
/// The text to encode.
/// The maximum number of tokens to encode.
/// If the tokenizer's normalization is enabled or is , this will be set to in its normalized form; otherwise, this value will be set to .
/// The token count can be generated which should be smaller than the maximum token count.
/// Indicate whether to consider pre-tokenization before tokenization.
/// Indicate whether to consider normalization before tokenization.
///
/// The index of the maximum encoding capacity within the processed text without surpassing the token limit.
/// It represents the index immediately following the last character to be included. In cases where no tokens fit, the result will be 0; conversely,
/// if all tokens fit, the result will be length of the input text or the if the normalization is enabled.
///
public int GetIndexByTokenCount(ReadOnlySpan text, int maxTokenCount, out string? normalizedText, out int tokenCount, bool considerPreTokenization = true, bool considerNormalization = true)
=> GetIndexByTokenCount(
null,
text,
new EncodeSettings { ConsiderPreTokenization = considerPreTokenization, ConsiderNormalization = considerNormalization, MaxTokenCount = maxTokenCount },
fromEnd: false,
out normalizedText,
out tokenCount);
///
/// Find the index of the maximum encoding capacity without surpassing the token limit.
///
/// The text to encode.
/// The maximum number of tokens to encode.
/// If the tokenizer's normalization is enabled or is , this will be set to in its normalized form; otherwise, this value will be set to .
/// The token count can be generated which should be smaller than the maximum token count.
/// Indicate whether to consider pre-tokenization before tokenization.
/// Indicate whether to consider normalization before tokenization.
///
/// The index of the maximum encoding capacity within the processed text without surpassing the token limit.
/// It represents the index of the first character to be included. In cases where no tokens fit, the result will be the text length; conversely,
/// if all tokens fit, the result will be zero.
///
public int GetIndexByTokenCountFromEnd(string text, int maxTokenCount, out string? normalizedText, out int tokenCount, bool considerPreTokenization = true, bool considerNormalization = true)
=> GetIndexByTokenCount(
text,
text.AsSpan(),
new EncodeSettings { ConsiderPreTokenization = considerPreTokenization, ConsiderNormalization = considerNormalization, MaxTokenCount = maxTokenCount },
fromEnd: true,
out normalizedText,
out tokenCount);
///
/// Find the index of the maximum encoding capacity without surpassing the token limit.
///
/// The text to encode.
/// The maximum number of tokens to encode.
/// If the tokenizer's normalization is enabled or is , this will be set to in its normalized form; otherwise, this value will be set to .
/// The token count can be generated which should be smaller than the maximum token count.
/// Indicate whether to consider pre-tokenization before tokenization.
/// Indicate whether to consider normalization before tokenization.
///
/// The index of the maximum encoding capacity within the processed text without surpassing the token limit.
/// It represents the index of the first character to be included. In cases where no tokens fit, the result will be the text length; conversely,
/// if all tokens fit, the result will be zero.
///
public int GetIndexByTokenCountFromEnd(ReadOnlySpan text, int maxTokenCount, out string? normalizedText, out int tokenCount, bool considerPreTokenization = true, bool considerNormalization = true)
=> GetIndexByTokenCount(
null,
text,
new EncodeSettings { ConsiderPreTokenization = considerPreTokenization, ConsiderNormalization = considerNormalization, MaxTokenCount = maxTokenCount },
fromEnd: true,
out normalizedText,
out tokenCount);
///
/// Decode the given ids, back to a String.
///
/// The list of ids that we want to decode.
/// The decoded string.
/// is null.
/// contains invalid data.
///
/// Types derived from may override this implementation to provide a more efficient implementation.
/// By default, it uses .
///
public virtual string Decode(IEnumerable ids)
{
if (ids is null)
{
throw new ArgumentNullException(nameof(ids));
}
int idCount = 0;
if (ids is ICollection c)
{
idCount = c.Count;
if (idCount == 0)
{
return string.Empty;
}
}
char[] destination = ArrayPool.Shared.Rent(
#if DEBUG
1); // to help validate growth logic
#else
idCount == 0 ? 1024 : idCount * 8); // arbitrary starting point / heuristic
#endif
while (true)
{
switch (Decode(ids, destination, out int idsConsumed, out int charsWritten))
{
case OperationStatus.Done:
string result = destination.AsSpan(0, charsWritten).ToString();
ArrayPool.Shared.Return(destination);
return result;
case OperationStatus.DestinationTooSmall:
long newSize = (long)destination.Length * 2;
if (newSize > int.MaxValue)
{
newSize = (long)destination.Length + 1;
if (newSize > int.MaxValue)
{
throw new OutOfMemoryException();
}
}
ArrayPool.Shared.Return(destination);
destination = ArrayPool.Shared.Rent((int)newSize);
break;
default:
throw new InvalidOperationException("The provided token IDs could not be decoded.");
}
}
}
///
/// Decode the given ids back to text and store the result in the span.
///
/// The list of ids that we want to decode.
/// The span to store the decoded text.
/// The number of ids consumed during the decoding.
/// The number of characters written to the destination span.
/// The operation status indicates whether all IDs were successfully decoded or if the is too small to contain the entire decoded result.
public abstract OperationStatus Decode(IEnumerable ids, Span destination, out int idsConsumed, out int charsWritten);
internal static IEnumerable<(int Offset, int Length)>? InitializeForEncoding(
string? text,
ReadOnlySpan textSpan,
bool considerPreTokenization,
bool considerNormalization,
Normalizer? normalizer,
PreTokenizer? preTokenizer,
out string? normalizedText,
out ReadOnlySpan textSpanToEncode,
out int fullTextLength)
{
normalizedText = null;
IEnumerable<(int Offset, int Length)>? splits = null;
if (text is null)
{
if (considerNormalization && (normalizer is not null))
{
normalizedText = normalizer.Normalize(textSpan.ToString());
textSpanToEncode = normalizedText.AsSpan();
fullTextLength = normalizedText.Length;
if (considerPreTokenization && preTokenizer is not null)
{
splits = preTokenizer.PreTokenize(normalizedText);
}
}
else
{
textSpanToEncode = textSpan;
fullTextLength = textSpan.Length;
if (considerPreTokenization && preTokenizer is not null)
{
splits = preTokenizer.PreTokenize(textSpan);
}
}
}
else
{
if (considerNormalization && (normalizer is not null))
{
normalizedText = normalizer.Normalize(text);
textSpanToEncode = normalizedText.AsSpan();
fullTextLength = normalizedText.Length;
if (considerPreTokenization && preTokenizer is not null)
{
splits = preTokenizer.PreTokenize(normalizedText);
}
}
else
{
textSpanToEncode = text.AsSpan();
fullTextLength = text.Length;
if (considerPreTokenization && preTokenizer is not null)
{
splits = preTokenizer.PreTokenize(text);
}
}
}
return splits;
}
}
}