AI 基础知识十七 BBPE字节级分词
引言
BBPE是BPE的升级版,BPE最初是一种数据压缩算法,用于通过迭代合并高频字节对来减少数据大小,BPE被引入自然语言处理(NLP)领域,作为分词(Tokenization)的核心方法。BPE 指的是字节对编码。英文单词大概有60万,中文汉字大概 有6万,如果都加到词表中,词表巨大、显存爆炸、训练极度不稳定,有部分单词很少用到白白占位,还有未登录词问题(OOV),如果用基本字节组合又太小,比如用26个字母 组合来表示一个单词,又太细了,字符序列变得很长。词表不能太大也不能太小,折中方案定制大小的词表,根据实际需求来设置词表大小。
BPE算法过程
1.拆分单词: 语料中每个单词拆分为字符序列
2.统计频率:计算所有相邻字符对的出现频率(出现次数)
3.合并高频对: 选择频率最高的字符对,合并为新子词,更新词表
4.迭代合并:重复步骤2-4,直到词表大小达到V或最高频对频率为1。
举个例子:"用电电电鳗电鳗会不会被电电死" 过程:
1. -> [用,电,电,电,鳗,电,鳗,会,不,会,被,电,电,死]
2 -> [ 用电, 电电,电鳗,鳗电,电鳗,鳗会,会不,不会,会被,被电,电电,电死]
统计 -> [用电:1, 电电: 2, 电鳗:2,鳗电:1,鳗会:1,会不:1,不会:1 ,会被:1, 电死:1], “电电”频率是2是当前最高,合并“电电”为新子词
3--> [用,电电,电,鳗,电,鳗,会,不,会,被, 电电, 死]
4--> [用电电, 电电电, 电鳗, 鳗电,电鳗, 鳗会,会不,不会,会被, 被电电, 电电死]
统计 -> [用电电: 1, 电电电: 1, 电鳗:2, 鳗电:1, 鳗会:1,会不:1,不会:1,会被:1, 被电电:1, 电电死:1]
“电鳗”频率是2,合并为新子词
5--> [用,电电,电鳗,电鳗,会,不,会,被, 电电, 死]
"用电电电鳗电鳗会不会被电电死"最后分词结果是 “用 电电 电鳗 电鳗 会 不 会 被 电电 死”
假设词表大小设置为500, 基础表是256个, 上面语料中只提取出两个词 “电电” “电鳗” 增加两个,id从0开始,“电电” :256 ,“电鳗” :257,对于 “用“ 字使用UTF8进行编码 三个字节表示。
UTF8编码
英文可以能过空格进行分词,如果有中英混合句子或其他语言,为了能跨语言统一我们使用UTF8进行编码,由初始词表 256 字节,加合并新词表构成。
比如中文一个汉字占 三个字节,第一个字节给出占位长度
int XBBPE::GetWordSize(uint8_t ch)
{
int len = 1;
if ((ch & 0x80) == 0)
{
len = 1; // ASCII
}
else if ((ch & 0xE0) == 0xC0)
{
len = 2;
}
else if ((ch & 0xF0) == 0xE0)
{
len = 3; // 中文
}
else if ((ch & 0xF8) == 0xF0)
{
len = 4;
}
else
{
len = 1;
}
return len;
}
VS默认用GBK编码 ,所以对输入文本要先换成UTF8编码
string XBBPE::ToUTF8(const string& strGbk)
{
return MultiByteToMultiByte(strGbk, CP_ACP, CP_UTF8);
}
string XBBPE::ToGBK(const string& strUtf8)
{
return MultiByteToMultiByte(strUtf8, CP_UTF8, CP_ACP);
}
string XBBPE::MultiByteToMultiByte(const string& str, UINT from, UINT bto)
{
int wide_size = MultiByteToWideChar(from, 0, str.c_str(), -1, NULL, 0);
std::wstring wideStr(wide_size, 0);
MultiByteToWideChar(from, 0, str.c_str(), -1, wideStr.data(), wide_size);
int utf8_size = WideCharToMultiByte(bto, 0, wideStr.data(), -1, NULL, 0, NULL, NULL);
std::string multiStr(utf8_size, 0);
WideCharToMultiByte(bto, 0, wideStr.data(), -1, multiStr.data(), utf8_size, NULL, NULL);
multiStr.pop_back();
return multiStr;
}
实现BBPE算法
实现思路
合并表由于多个字合成 每字占用3个字节 (中文)

struct WordIdKey
{
WordIdKey()
{
}
WordIdKey(const VectorUint8& key)
{
len = key.size();
memcpy(idKey, key.data(), min(key.size(), sizeof(idKey)));
}
bool Append(const WordIdKey& k)
{
bool b = false;
if (len + k.len < sizeof(idKey))
{
memcpy(idKey + len, k.idKey,k.len);
len += k.len;
b = true;
}
return b;
}
WordIdKey(const string& key)
{
len = key.size();
memcpy(idKey, key.data(), min(key.size(), sizeof(idKey)));
}
uint8_t idKey[MaxKeyCount] ={ 0 };
size_t len = 0;
bool operator == (const WordIdKey& other) const
{
return len == other.len && memcmp(idKey, other.idKey, sizeof(idKey)) == 0;
}
bool operator < (const WordIdKey& other) const
{
if (len != other.len)
{
return len < other.len;
}
return memcmp(idKey, other.idKey, sizeof(idKey)) < 0;
}
void operator = (const WordIdKey& other)
{
if (this == &other)
{
return ;
}
len = other.len;
memcpy(idKey, other.idKey, sizeof(idKey));
}
};
typedef unordered_map<WordIdKey, int64_t> MapEncoderWordList; ///
typedef unordered_map<int64_t, WordIdKey> MapDecoderWordList; ///
MapEncoderWordList 扫描文本生成int64编码词表,MapDecoderWordList是编码转文本
1.构成初始词表
256 字节基础符号追 加 中文标点符号(,。),统计中文时不带标点符号
void XBBPE::InitData(void)
{
m_mapEncoderList.clear();
m_mapDecoderList.clear();
for (int i = 0; i < 256; i++)
{
VectorUint8 b;
b.push_back(i);
AddNewKeyToWordList(b);
}
std::vector<VectorUint8> filterSyms =
{
{0xC2, 0xB7},
{0xEF, 0xBC, 0x8C},
{0xEF, 0xBC, 0x9F},
{0xEF, 0xBC, 0x81},
{0xE3, 0x80, 0x82},
{0xE3, 0x80, 0x80},
{0xE2, 0x80, 0x8B}
};
for (auto& f : filterSyms)
{
AddNewKeyToWordList(f);
}
}
2.预处理文本
按标点符号空格划分文本,
string strReg = R"(\x0A|\x3F|\x20|\x21|\x22|\x2C|\x2E|\xC2\xB7|\xEF\xBC\x8C|\xEF\xBC\x9F|\xEF\xBC\x81|\xE3\x80\x82|\xE3\x80\x80|\xE2\x80\x8B)";
auto special = regex(strReg);
for (auto& slist : textList)
{
auto strText = ToUTF8(slist);
sregex_token_iterator it(strText.begin(), strText.end(), special, { -1,1 });
sregex_token_iterator end;
for (auto seq = it; seq != end; seq++)
{
}
}
3.统计频率找出最高词频
for (auto& list : v2WordList)
{
for (size_t i = 0; i+1 < list.size(); i++)
{
WordIdKey m;
m.Append(list[i]);
m.Append(list[i+1]);
if (merge.find(m) == merge.end())
{
merge.emplace(m, 0);
}
else
{
merge[m]++;
}
if (maxPair < merge[m])
{
maxPair = merge[m];
maxWord = m;
}
}
}
4.编码
void XBBPE::Encode(const string& textGbk, VectorInt64& ids)
{
ids.clear();
auto special = regex(R"(<[^>]*>)");
auto text = ToUTF8(textGbk);
sregex_token_iterator it(text.begin(), text.end(), special, { -1, 0 });
sregex_token_iterator end;
for (auto seq = it; seq != end; ++seq)
{
string s = *seq;
if (s.empty())
{
continue;
}
VectorWord vWordList;
ToTextVectorWord(s, vWordList);
for (size_t i = 0; i < vWordList.size(); i++)
{
WordIdKey m(vWordList[i]);
do
{
WordIdKey m2 = m;
if (i+1 < vWordList.size())
{
m2.Append(vWordList[i + 1]);
}
if (!IsInWordList(m2))
{
break;
}
m = m2;
i += 1;
} while (i + 1 < vWordList.size());
GetWordEncode(m, ids);
}
}
}
void XBBPE::ToTextVectorWord(const string& strUtf8, VectorWord& vWordList)
{
vWordList.clear();
for (size_t i = 0; i < strUtf8.size(); i++)
{
int len = GetWordSize(strUtf8[i]);
VectorUint8 word;
for (int j = 0; j < len; j++)
{
word.push_back(strUtf8[i + j]);
}
i += len - 1;
vWordList.push_back({ word });
}
}
void XBBPE::GetWordEncode(WordIdKey& word, VectorInt64& vList)
{
if (m_mapEncoderList.find(word) != m_mapEncoderList.end())
{
vList.push_back(m_mapEncoderList.at(word));
}
else
{
for (int i = 0; i < word.len; i++)
{
WordIdKey tm;
tm.idKey[0] = word.idKey[i];
tm.len = 1;
vList.push_back(m_mapEncoderList.at(tm));
}
}
}
5. 解码
string XBBPE::Decoded(const VectorInt64& ids)
{
VectorUint8 vList;
for (auto& id : ids)
{
auto word = m_mapDecoderList.at(id);
for (int i = 0;i < word.len; i++)
{
vList.push_back(word.idKey[i]);
}
}
string str(vList.begin(), vList.end());
str = ToGBK(str);
return str;
}
c++实现BBPE完整代码
XBBPE.h
#pragma once
#define VocabSize 800
#define BBPE_PATH "xBBPE.bin"
#define MaxKeyCount (3*8) //
#define BOS "<S>"
#define EOS "</S>"
#define PAD "<P>"
typedef vector<string> VectorString;
typedef vector<uint8_t> VectorUint8;
struct WordIdKey
{
WordIdKey()
{
}
WordIdKey(const VectorUint8& key)
{
len = key.size();
memcpy(idKey, key.data(), min(key.size(), sizeof(idKey)));
}
bool Append(const WordIdKey& k)
{
bool b = false;
if (len + k.len < sizeof(idKey))
{
memcpy(idKey + len, k.idKey,k.len);
len += k.len;
b = true;
}
return b;
}
WordIdKey(const string& key)
{
len = key.size();
memcpy(idKey, key.data(), min(key.size(), sizeof(idKey)));
}
uint8_t idKey[MaxKeyCount] ={ 0 };
size_t len = 0;
bool operator == (const WordIdKey& other) const
{
return len == other.len && memcmp(idKey, other.idKey, sizeof(idKey)) == 0;
}
bool operator < (const WordIdKey& other) const
{
if (len != other.len)
{
return len < other.len;
}
return memcmp(idKey, other.idKey, sizeof(idKey)) < 0;
}
void operator = (const WordIdKey& other)
{
if (this == &other)
{
return ;
}
len = other.len;
memcpy(idKey, other.idKey, sizeof(idKey));
}
};
template<>
struct std::hash<WordIdKey>
{
size_t operator()(const WordIdKey& k) const noexcept
{
size_t h = 0;
for (int i = 0; i < k.len; ++i)
{
h ^= std::hash<uint8_t>{}(k.idKey[i]) + 0x9e3779b9 + (h << 6) + (h >> 2);
}
return h;
}
};
typedef vector<int64_t> VectorInt64;
typedef vector<WordIdKey> VectorWord;
typedef vector<VectorWord> Vector2Word;
typedef unordered_map<WordIdKey, int64_t> MapEncoderWordList; ///
typedef unordered_map<int64_t, WordIdKey> MapDecoderWordList; ///
typedef map<int64_t, WordIdKey> MapSingleWord;
//typedef map<WordIdKey, int64_t> MapPairWordCount;
//typedef vector<pair<size_t, size_t>> VectorPairWordIndex;
string GetOutputPath();
class XBBPE
{
public:
XBBPE();
~XBBPE();
void LoadDataFileTrain(const string& paths, uint32_t vocabSize = VocabSize);
void Encode(const string& text, VectorInt64& ids);
string Decoded(const VectorInt64& ids);
int64_t GetBOS();
int64_t GetEOS();
int64_t GetPAD();
int64_t GetDictionaryCount()
{
return m_mapEncoderList.size();
}
VectorTrainEncoded& GetTrainData()
{
return m_vectorTrainEncoded;
}
private:
void Train(const VectorString& textList, uint32_t vocabSize);
void InitData(void);
bool LoadFile(const string& path = BBPE_PATH);
void SaveFile(const string& path = BBPE_PATH);
int GetWordSize(uint8_t ch);
bool IsInWordList(const WordIdKey& key);
void AddNewKeyToWordList(const VectorUint8& vKey);
void AddNewKeyToWordList(const WordIdKey& key);
string ToGBK(const string& strUtf8);
string ToUTF8(const string& strGbk);
string MultiByteToMultiByte(const string& str, UINT from, UINT bto);
WordIdKey& MergeMaxPairWord(Vector2Word& v2WordList, MapSingleWord& vSingleWordList, bool del);
void AddSpecialTokens(const VectorString& tokens);
void ToTextVectorWord(const string& strUtf8, VectorWord& vWordList);
void GetWordEncode(WordIdKey& word, VectorInt64& vList);
private:
MapEncoderWordList m_mapEncoderList;
MapDecoderWordList m_mapDecoderList;
VectorTrainText m_vectorTrainText;
VectorTrainEncoded m_vectorTrainEncoded;
};
XBBPE.cpp
#include "pch.h"
#include "XBBPE.h"
string GetOutputPath()
{
auto s = std::filesystem::current_path().string();
size_t len = s.size() - s.find_last_of("\\");
s.erase(s.end() - len, s.end());
s += "\\tmpbin\\";
if (!filesystem::exists(s))
{
filesystem::create_directories(s);
}
return s;
}
XBBPE::XBBPE()
{
/*
VectorString corpus =
{
"用电电电鳗电鳗会不会被电电死?",
"bbpe 是 byte level bpe 分词算法。",
"bpe 算法用于大模型 token 编码。",
"bbpe 基于 utf8 字节合并中文英文。",
"token 编码电鳗放电测试。",
"token to a ab abc abc abcd abcf.,,。"
};
Train(corpus, 1000);
*/
if (!LoadFile())
{
InitData();
LoadDataFileTrain("tangshi.data.txt");
// LoadDataFileTrain("HelloWorld.txt");
}
/*
for (auto& item: m_vectorTrainEncoded)
{
string a = Decoded(item);
cout << a << endl;
}
*/
}
XBBPE::~XBBPE()
{
}
void XBBPE::LoadDataFileTrain(const string& paths, uint32_t vocabSize)
{
m_vectorTrainText.clear();
auto xPath = GetOutputPath() + paths;
std::ifstream ifs(xPath);
bool bopen = ifs.is_open();
std::stringstream ss;
ss << ifs.rdbuf();
ifs.close();
std::string line;
while (true)
{
while (getline(ss, line) && line.empty());
if (!ss)
{
break;
}
m_vectorTrainText.push_back(line);
while (getline(ss, line))
{
if (line.empty())
{
break;
}
m_vectorTrainText.push_back(line);
}
m_vectorTrainText[m_vectorTrainText.size() - 1] += "\n";
}
VectorString vstring;
for(auto& v : m_vectorTrainText)
{
vstring.push_back(v);
}
Train(vstring, vocabSize);
for (auto& v : m_vectorTrainText)
{
VectorInt64 item;
Encode(v, item);
m_vectorTrainEncoded.push_back(item);
}
SaveFile();
}
bool XBBPE::LoadFile(const string& path)
{
auto binPath = GetOutputPath() + path;
ifstream infs(binPath, ios::binary);
if (!infs.is_open())
{
return false;
}
m_vectorTrainEncoded.clear();
m_mapEncoderList.clear();
m_mapDecoderList.clear();
size_t count = 0;
infs.read((char*)&count, sizeof(size_t));
for (size_t i = 0; i < count; i++)
{
pair<WordIdKey, int64_t> item;
infs.read((char*)&item, sizeof(item));
m_mapEncoderList.emplace(item);
m_mapDecoderList.emplace(item.second, item.first);
}
count = 0;
infs.read((char*)&count, sizeof(size_t));
for (int i = 0; i < count; i++)
{
size_t len = 0;
infs.read((char*)&len, sizeof(size_t));
VectorInt64 item;
item.resize(len);
infs.read((char*)item.data(), len * sizeof(int64_t));
m_vectorTrainEncoded.push_back(item);
}
infs.close();
return !m_mapEncoderList.empty();
}
void XBBPE::SaveFile(const string& path)
{
auto binPath = GetOutputPath() + path;
remove(binPath.c_str());
ofstream outfs(binPath, ios::binary);
size_t count = m_mapEncoderList.size();
outfs.write((const char* ) & count, sizeof(count));
for (const auto& pair : m_mapEncoderList)
{
outfs.write((const char*)&pair, sizeof(pair));
}
count = m_vectorTrainEncoded.size();
outfs.write((const char*)&count, sizeof(count));
for (auto& v : m_vectorTrainEncoded)
{
count = v.size();
outfs.write((const char*)&count, sizeof(count));
outfs.write((const char*)v.data(), count * sizeof(int64_t));
}
outfs.close();
}
void XBBPE::InitData(void)
{
m_mapEncoderList.clear();
m_mapDecoderList.clear();
for (int i = 0; i < 256; i++)
{
VectorUint8 b;
b.push_back(i);
AddNewKeyToWordList(b);
}
std::vector<VectorUint8> filterSyms =
{
{0xC2, 0xB7},
{0xEF, 0xBC, 0x8C},
{0xEF, 0xBC, 0x9F},
{0xEF, 0xBC, 0x81},
{0xE3, 0x80, 0x82},
{0xE3, 0x80, 0x80},
{0xE2, 0x80, 0x8B}
};
for (auto& f : filterSyms)
{
AddNewKeyToWordList(f);
}
}
int XBBPE::GetWordSize(uint8_t ch)
{
int len = 1;
if ((ch & 0x80) == 0)
{
len = 1; // ASCII
}
else if ((ch & 0xE0) == 0xC0)
{
len = 2;
}
else if ((ch & 0xF0) == 0xE0)
{
len = 3; // 中文
}
else if ((ch & 0xF8) == 0xF0)
{
len = 4;
}
else
{
len = 1;
}
return len;
}
bool XBBPE::IsInWordList(const WordIdKey& key)
{
return m_mapEncoderList.find(key) != m_mapEncoderList.end();
}
void XBBPE::AddNewKeyToWordList(const VectorUint8& vKey)
{
string key(vKey.begin(), vKey.end());
AddNewKeyToWordList({ key });
}
void XBBPE::AddNewKeyToWordList(const WordIdKey& key)
{
auto id = m_mapEncoderList.size();
if (!IsInWordList(key))
{
m_mapEncoderList.emplace(key, id);
m_mapDecoderList.emplace(id, key);
}
}
string XBBPE::ToUTF8(const string& strGbk)
{
return MultiByteToMultiByte(strGbk, CP_ACP, CP_UTF8);
}
string XBBPE::ToGBK(const string& strUtf8)
{
return MultiByteToMultiByte(strUtf8, CP_UTF8, CP_ACP);
}
string XBBPE::MultiByteToMultiByte(const string& str, UINT from, UINT bto)
{
int wide_size = MultiByteToWideChar(from, 0, str.c_str(), -1, NULL, 0);
std::wstring wideStr(wide_size, 0);
MultiByteToWideChar(from, 0, str.c_str(), -1, wideStr.data(), wide_size);
int utf8_size = WideCharToMultiByte(bto, 0, wideStr.data(), -1, NULL, 0, NULL, NULL);
std::string multiStr(utf8_size, 0);
WideCharToMultiByte(bto, 0, wideStr.data(), -1, multiStr.data(), utf8_size, NULL, NULL);
multiStr.pop_back();
return multiStr;
}
void XBBPE::Train(const VectorString& textList, uint32_t vocabSize)
{
vocabSize = vocabSize - 7;
string strReg = R"(\x0A|\x3F|\x20|\x21|\x22|\x2C|\x2E|\xC2\xB7|\xEF\xBC\x8C|\xEF\xBC\x9F|\xEF\xBC\x81|\xE3\x80\x82|\xE3\x80\x80|\xE2\x80\x8B)";
auto special = regex(strReg);
InitData();
WordIdKey key;
Vector2Word v2WordList;
for (auto& slist : textList)
{
auto strText = ToUTF8(slist);
sregex_token_iterator it(strText.begin(), strText.end(), special, { -1,1 });
sregex_token_iterator end;
for (auto seq = it; seq != end; seq++)
{
VectorWord vItem;
string s = *seq;
if (s.empty())
{
continue;
}
//cout << ToGBK(s) << endl;
for (int i = 0; i < s.size(); i++)
{
int len = GetWordSize(s[i]);
VectorUint8 word;
for (int j = 0; j < len; j++)
{
word.push_back(s[i + j]);
}
i += len - 1;
vItem.push_back({ word });
}
if (1 < vItem.size())
{
v2WordList.push_back(vItem);
}
}
}
MapSingleWord vSingleWordList;
bool del = true;
while (m_mapEncoderList.size() < vocabSize)
{
WordIdKey addKey = MergeMaxPairWord(v2WordList, vSingleWordList,del);
if (del)
{
for (auto it = vSingleWordList.rbegin(); it != vSingleWordList.rend(); ++it)
{
if (m_mapEncoderList.size() < vocabSize)
{
AddNewKeyToWordList(it->second);
//string kk((char*)it->second.idKey);
//cout << ToGBK(kk) << endl;
}
else
{
break;
}
}
}
del = false;
if (addKey.len == 0 || vocabSize <= m_mapEncoderList.size() )
{
break;
}
//string kk((char*)addKey.idKey);
//cout << ToGBK(kk) << endl;
AddNewKeyToWordList(addKey);
}
VectorString sp;
sp.push_back(PAD);
sp.push_back(BOS);
sp.push_back(EOS);
AddSpecialTokens(sp);
}
void XBBPE::AddSpecialTokens(const VectorString& tokens)
{
for (auto& slist : tokens)
{
auto strutf8 = ToUTF8(slist);
for (int i=0; i < strutf8.size(); i++)
{
VectorUint8 item(strutf8.begin(), strutf8.begin()+1+i);
AddNewKeyToWordList(item);
}
}
}
WordIdKey& XBBPE::MergeMaxPairWord(Vector2Word& v2WordList, MapSingleWord& vSingleWordList, bool del)
{
MapEncoderWordList single;
MapEncoderWordList merge;
VectorWord maxlist;
size_t maxPair = 0;
WordIdKey maxWord;
for (auto& list : v2WordList)
{
for (size_t i = 0; i+1 < list.size(); i++)
{
WordIdKey m;
m.Append(list[i]);
m.Append(list[i+1]);
if (merge.find(m) == merge.end())
{
merge.emplace(m, 0);
}
else
{
merge[m]++;
}
if (maxPair < merge[m])
{
maxPair = merge[m];
maxWord = m;
}
}
}
for (auto& list : v2WordList)
{
for (int i = 0; i + 1 < list.size(); i++)
{
WordIdKey m;
if (1 <= i)
{
m.Append(list[i - 1]);
m.Append(list[i]);
}
WordIdKey m2;
m2.Append(list[i]);
m2.Append(list[i + 1]);
if (merge[m] == 0 && merge[m2] == 0 && del)
{
if (single.find(list[i]) == single.end())
{
single.emplace(list[i], 0);
}
else
{
single[list[i]]++;
}
list.erase(list.begin()+i);
i--;
}
else if (maxWord == m2 && merge[m2] == maxPair)
{
list[i] = maxWord;
list.erase(list.begin() + i + 1);
}
}
}
if (del)
{
int64_t baseMax = single.size()*10;
int64_t index = 0;
for (auto& key : single)
{
if (3 <= key.first.len)
{
int64_t keyId = key.second * baseMax + index;
index++;
vSingleWordList.emplace(keyId, key.first);
}
}
v2WordList.erase(std::remove_if(v2WordList.begin(), v2WordList.end(), [&](const VectorWord& vw)
{
bool b = vw.size() <= 1;
return b;
}), v2WordList.end());
}
return maxWord;
}
int64_t XBBPE::GetBOS()
{
int64_t id = 0;
VectorInt64 ids;
Encode(BOS, ids);
if (0 < ids.size())
{
id = ids.at(0);
}
return id;
}
int64_t XBBPE::GetEOS()
{
int64_t id = 0;
VectorInt64 ids;
Encode(EOS, ids);
if (0 < ids.size())
{
id = ids.at(0);
}
return id;
}
int64_t XBBPE::GetPAD()
{
int64_t id = 0;
VectorInt64 ids;
Encode(PAD, ids);
if (0 < ids.size())
{
id = ids.at(0);
}
return id;
}
void XBBPE::Encode(const string& textGbk, VectorInt64& ids)
{
ids.clear();
auto special = regex(R"(<[^>]*>)");
auto text = ToUTF8(textGbk);
sregex_token_iterator it(text.begin(), text.end(), special, { -1, 0 });
sregex_token_iterator end;
for (auto seq = it; seq != end; ++seq)
{
string s = *seq;
if (s.empty())
{
continue;
}
VectorWord vWordList;
ToTextVectorWord(s, vWordList);
for (size_t i = 0; i < vWordList.size(); i++)
{
WordIdKey m(vWordList[i]);
do
{
WordIdKey m2 = m;
if (i+1 < vWordList.size())
{
m2.Append(vWordList[i + 1]);
}
if (!IsInWordList(m2))
{
break;
}
m = m2;
i += 1;
} while (i + 1 < vWordList.size());
GetWordEncode(m, ids);
}
}
}
void XBBPE::ToTextVectorWord(const string& strUtf8, VectorWord& vWordList)
{
vWordList.clear();
for (size_t i = 0; i < strUtf8.size(); i++)
{
int len = GetWordSize(strUtf8[i]);
VectorUint8 word;
for (int j = 0; j < len; j++)
{
word.push_back(strUtf8[i + j]);
}
i += len - 1;
vWordList.push_back({ word });
}
}
void XBBPE::GetWordEncode(WordIdKey& word, VectorInt64& vList)
{
if (m_mapEncoderList.find(word) != m_mapEncoderList.end())
{
vList.push_back(m_mapEncoderList.at(word));
}
else
{
for (int i = 0; i < word.len; i++)
{
WordIdKey tm;
tm.idKey[0] = word.idKey[i];
tm.len = 1;
vList.push_back(m_mapEncoderList.at(tm));
}
}
}
string XBBPE::Decoded(const VectorInt64& ids)
{
VectorUint8 vList;
for (auto& id : ids)
{
auto word = m_mapDecoderList.at(id);
for (int i = 0;i < word.len; i++)
{
vList.push_back(word.idKey[i]);
}
}
string str(vList.begin(), vList.end());
str = ToGBK(str);
return str;
}
BBPE概念和实现不是很难, 直接阅读代码就好了
感谢大家的支持。
AtomGit 是由开放原子开源基金会联合 CSDN 等生态伙伴共同推出的新一代开源与人工智能协作平台。平台坚持“开放、中立、公益”的理念,把代码托管、模型共享、数据集托管、智能体开发体验和算力服务整合在一起,为开发者提供从开发、训练到部署的一站式体验。
更多推荐



所有评论(0)