引言

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概念和实现不是很难, 直接阅读代码就好了 


感谢大家的支持。

Logo

AtomGit 是由开放原子开源基金会联合 CSDN 等生态伙伴共同推出的新一代开源与人工智能协作平台。平台坚持“开放、中立、公益”的理念,把代码托管、模型共享、数据集托管、智能体开发体验和算力服务整合在一起,为开发者提供从开发、训练到部署的一站式体验。

更多推荐