feat: Initial Commit
This commit is contained in:
@@ -0,0 +1,232 @@
|
||||
import json
|
||||
from collections import Counter
|
||||
from pathlib import Path
|
||||
|
||||
from config import (
|
||||
PAD_TOKEN,
|
||||
START_TOKEN,
|
||||
END_TOKEN,
|
||||
UNK_TOKEN,
|
||||
SPECIAL_TOKENS,
|
||||
)
|
||||
|
||||
|
||||
class Tokenizer:
|
||||
def __init__(self):
|
||||
self.text_to_id = {}
|
||||
self.id_to_text = {}
|
||||
|
||||
self.output_to_id = {}
|
||||
self.id_to_output = {}
|
||||
|
||||
# =====================================================
|
||||
# Build Vocabulary
|
||||
# =====================================================
|
||||
|
||||
def build(self, dataset_path):
|
||||
|
||||
dataset_path = Path(dataset_path)
|
||||
|
||||
with open(dataset_path, "r", encoding="utf8") as f:
|
||||
data = json.load(f)
|
||||
|
||||
text_counter = Counter()
|
||||
output_counter = Counter()
|
||||
|
||||
for sample in data:
|
||||
|
||||
# ---------------------------
|
||||
# Prompt tokens
|
||||
# ---------------------------
|
||||
|
||||
prompt = sample["prompt"].lower().split()
|
||||
|
||||
text_counter.update(prompt)
|
||||
|
||||
# ---------------------------
|
||||
# Output tokens
|
||||
# ---------------------------
|
||||
|
||||
for block in sample["blocks"]:
|
||||
|
||||
x = block["x"]
|
||||
y = block["y"]
|
||||
z = block["z"]
|
||||
b = block["block"]
|
||||
|
||||
output_counter.update([
|
||||
f"X_{x}",
|
||||
f"Y_{y}",
|
||||
f"Z_{z}",
|
||||
f"BLOCK_{b}"
|
||||
])
|
||||
|
||||
# Special tokens
|
||||
|
||||
text_vocab = SPECIAL_TOKENS + sorted(text_counter.keys())
|
||||
output_vocab = SPECIAL_TOKENS + sorted(output_counter.keys())
|
||||
|
||||
self.text_to_id = {
|
||||
token: i for i, token in enumerate(text_vocab)
|
||||
}
|
||||
|
||||
self.id_to_text = {
|
||||
i: token for token, i in self.text_to_id.items()
|
||||
}
|
||||
|
||||
self.output_to_id = {
|
||||
token: i for i, token in enumerate(output_vocab)
|
||||
}
|
||||
|
||||
self.id_to_output = {
|
||||
i: token for token, i in self.output_to_id.items()
|
||||
}
|
||||
|
||||
# =====================================================
|
||||
# Prompt Encoding
|
||||
# =====================================================
|
||||
|
||||
def encode_prompt(self, prompt):
|
||||
|
||||
tokens = prompt.lower().split()
|
||||
|
||||
ids = [
|
||||
self.text_to_id[START_TOKEN]
|
||||
]
|
||||
|
||||
for token in tokens:
|
||||
|
||||
ids.append(
|
||||
self.text_to_id.get(
|
||||
token,
|
||||
self.text_to_id[UNK_TOKEN]
|
||||
)
|
||||
)
|
||||
|
||||
ids.append(
|
||||
self.text_to_id[END_TOKEN]
|
||||
)
|
||||
|
||||
return ids
|
||||
|
||||
def decode_prompt(self, ids):
|
||||
|
||||
words = []
|
||||
|
||||
for idx in ids:
|
||||
|
||||
token = self.id_to_text[idx]
|
||||
|
||||
if token in SPECIAL_TOKENS:
|
||||
continue
|
||||
|
||||
words.append(token)
|
||||
|
||||
return " ".join(words)
|
||||
|
||||
# =====================================================
|
||||
# Structure Encoding
|
||||
# =====================================================
|
||||
|
||||
def encode_blocks(self, blocks):
|
||||
|
||||
ids = [
|
||||
self.output_to_id[START_TOKEN]
|
||||
]
|
||||
|
||||
for block in blocks:
|
||||
|
||||
ids.extend([
|
||||
self.output_to_id[f"X_{block['x']}"],
|
||||
self.output_to_id[f"Y_{block['y']}"],
|
||||
self.output_to_id[f"Z_{block['z']}"],
|
||||
self.output_to_id[f"BLOCK_{block['block']}"],
|
||||
])
|
||||
|
||||
ids.append(
|
||||
self.output_to_id[END_TOKEN]
|
||||
)
|
||||
|
||||
return ids
|
||||
|
||||
def decode_blocks(self, ids):
|
||||
|
||||
tokens = []
|
||||
|
||||
for idx in ids:
|
||||
|
||||
token = self.id_to_output[idx]
|
||||
|
||||
if token in SPECIAL_TOKENS:
|
||||
continue
|
||||
|
||||
tokens.append(token)
|
||||
|
||||
blocks = []
|
||||
|
||||
i = 0
|
||||
|
||||
while i + 3 < len(tokens):
|
||||
|
||||
x = int(tokens[i][2:])
|
||||
y = int(tokens[i + 1][2:])
|
||||
z = int(tokens[i + 2][2:])
|
||||
block = tokens[i + 3][6:]
|
||||
|
||||
blocks.append({
|
||||
"x": x,
|
||||
"y": y,
|
||||
"z": z,
|
||||
"block": block
|
||||
})
|
||||
|
||||
i += 4
|
||||
|
||||
return blocks
|
||||
|
||||
# =====================================================
|
||||
# Save / Load
|
||||
# =====================================================
|
||||
|
||||
def save(self, folder):
|
||||
|
||||
folder = Path(folder)
|
||||
folder.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
with open(folder / "text_vocab.json", "w") as f:
|
||||
json.dump(self.text_to_id, f, indent=4)
|
||||
|
||||
with open(folder / "output_vocab.json", "w") as f:
|
||||
json.dump(self.output_to_id, f, indent=4)
|
||||
|
||||
def load(self, folder):
|
||||
|
||||
folder = Path(folder)
|
||||
|
||||
with open(folder / "text_vocab.json", "r") as f:
|
||||
self.text_to_id = json.load(f)
|
||||
|
||||
with open(folder / "output_vocab.json", "r") as f:
|
||||
self.output_to_id = json.load(f)
|
||||
|
||||
self.id_to_text = {
|
||||
int(v): k
|
||||
for k, v in self.text_to_id.items()
|
||||
}
|
||||
|
||||
self.id_to_output = {
|
||||
int(v): k
|
||||
for k, v in self.output_to_id.items()
|
||||
}
|
||||
|
||||
# =====================================================
|
||||
# Properties
|
||||
# =====================================================
|
||||
|
||||
@property
|
||||
def text_vocab_size(self):
|
||||
return len(self.text_to_id)
|
||||
|
||||
@property
|
||||
def output_vocab_size(self):
|
||||
return len(self.output_to_id)
|
||||
Reference in New Issue
Block a user