[refactor] move rbm.label → label.label
Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
This commit is contained in:
@@ -0,0 +1,119 @@
|
||||
from rbm.matrix import Mat, np
|
||||
from rbm.model import Model
|
||||
from rbm.entity import Entity, EntityParams, TrainingParams
|
||||
|
||||
import math
|
||||
from enum import Enum
|
||||
|
||||
class Label:
|
||||
class EncodingType(Enum):
|
||||
Binary = 1
|
||||
OneHot = 2
|
||||
|
||||
class Fitter(Model):
|
||||
def __init__(self, dim: tuple[int,int], work_dir: str = '.'):
|
||||
super().__init__(f"label-{dim[0]}x{dim[1]}", work_dir)
|
||||
self.unit1 = Entity(dim, EntityParams(), TrainingParams(num_epochs=10000, do_rao_blackwell=False))
|
||||
|
||||
def forward(self, x: Mat) -> Mat:
|
||||
return self.unit1.forward(x)
|
||||
|
||||
def reconstruct(self, x: Mat):
|
||||
return self.unit1.reconstruct(x)
|
||||
|
||||
def __init__(self, voc_size, num_dim, encoding: EncodingType, work_dir):
|
||||
self.voc_size = voc_size
|
||||
if encoding == Label.EncodingType.Binary:
|
||||
self.label2vec = self.label2vec_binary
|
||||
self.vec2label = self.vec2label_binary
|
||||
self.fitter = Label.Fitter((round(math.log(voc_size,2)), num_dim), work_dir)
|
||||
elif encoding == Label.EncodingType.OneHot:
|
||||
self.label2vec = self.label2vec_onehot
|
||||
self.vec2label = self.vec2label_onehot
|
||||
self.fitter = Label.Fitter((voc_size, num_dim), work_dir)
|
||||
|
||||
def fit(self, list_of_labels: np.array):
|
||||
label_vecs = self.label2vec(list_of_labels)
|
||||
self.fitter.train(label_vecs)
|
||||
|
||||
def encode(self, list_of_labels: np.array):
|
||||
label_vecs = self.label2vec(list_of_labels)
|
||||
enc_data = self.fitter.forward(label_vecs)
|
||||
return enc_data, label_vecs
|
||||
|
||||
def decode(self, list_of_encoded_label: np.array):
|
||||
dec_data = self.fitter.reconstruct(list_of_encoded_label)
|
||||
return dec_data
|
||||
|
||||
def label2vec_binary(self, list_of_labels: np.array):
|
||||
num_bits = round(math.log(self.voc_size,2))
|
||||
batch = np.zeros((len(list_of_labels), num_bits))
|
||||
for index, label in enumerate(list_of_labels):
|
||||
b = self._enc_binary(label, num_bits)
|
||||
batch[index, :] = np.array(b)
|
||||
|
||||
return batch
|
||||
|
||||
def vec2label_binary(self, label_vecs: np.array, thresh=0.9):
|
||||
labels = np.zeros((len(label_vecs)))
|
||||
for index, label_vecs in enumerate(label_vecs):
|
||||
labels[index] = self._dec_binary(label_vecs > thresh)
|
||||
return labels
|
||||
|
||||
@staticmethod
|
||||
def _enc_binary(number: int, num_bits):
|
||||
res = [1 if s == '1' else 0 for s in format(number, f'{num_bits}b')]
|
||||
return res
|
||||
|
||||
@staticmethod
|
||||
def _dec_binary(vec: np.array):
|
||||
res = 0
|
||||
e = 1
|
||||
for v in np.flip(vec):
|
||||
res += v*e
|
||||
e *= 2
|
||||
|
||||
return res
|
||||
|
||||
def label2vec_onehot(self, list_of_labels: np.array):
|
||||
res = []
|
||||
for label in list_of_labels:
|
||||
vec = np.zeros(self.voc_size)
|
||||
vec[label] = 1.0
|
||||
res.append(vec)
|
||||
return np.stack(res, axis=0)
|
||||
|
||||
def vec2label_onehot(self, label_vecs: np.array, thresh=0.9):
|
||||
return np.argmax(label_vecs, axis=-1)
|
||||
|
||||
if __name__ == "__main__":
|
||||
voc_size = 100
|
||||
|
||||
# Create Labeler
|
||||
enc = Label(voc_size, 16, work_dir="../../results")
|
||||
|
||||
# Load
|
||||
enc.fitter.load()
|
||||
|
||||
# Train
|
||||
train_labels = list(range(voc_size))
|
||||
enc.fit(train_labels)
|
||||
|
||||
# Save fitter state
|
||||
enc.fitter.save()
|
||||
|
||||
# Encode test labels
|
||||
test_labels = np.array([1, 23, 99, 37, 55, 7, 31, 10, 19, 70])
|
||||
encoded, binary_labels = enc.encode(train_labels)
|
||||
|
||||
# Decode
|
||||
decoded_soft = enc.decode(encoded)
|
||||
decode_hard = (decoded_soft > 0.9).astype(int)
|
||||
|
||||
# Convert soft state into labels
|
||||
labels_reconst = enc.vec2label(decoded_soft)
|
||||
print(labels_reconst)
|
||||
|
||||
print(f"Soft Error: {np.sum((decoded_soft - binary_labels) ** 2)}")
|
||||
print(f"Hard Error: {np.sum((decode_hard - binary_labels) ** 2)}")
|
||||
print("Test: [passed]")
|
||||
Reference in New Issue
Block a user