tune label fitter: epochs, momentum, encoding, decode threshold
- reduce num_epochs to 5000 and add momentum=0.9 - switch to OneHot encoding in Label constructor - call fitter.init(0.1) before load for weight initialisation - use train_labels as test set, lower hard-decode threshold to 0.5 - improve print output Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
This commit is contained in:
+8
-6
@@ -13,7 +13,7 @@ class Label:
|
|||||||
class Fitter(Model):
|
class Fitter(Model):
|
||||||
def __init__(self, dim: tuple[int,int], work_dir: str = '.'):
|
def __init__(self, dim: tuple[int,int], work_dir: str = '.'):
|
||||||
super().__init__(f"label-{dim[0]}x{dim[1]}", work_dir)
|
super().__init__(f"label-{dim[0]}x{dim[1]}", work_dir)
|
||||||
self.unit1 = Entity(dim, EntityParams(), TrainingParams(num_epochs=10000, do_rao_blackwell=False))
|
self.unit1 = Entity(dim, EntityParams(), TrainingParams(num_epochs=5000, do_rao_blackwell=False, momentum=0.9))
|
||||||
|
|
||||||
def forward(self, x: Mat) -> Mat:
|
def forward(self, x: Mat) -> Mat:
|
||||||
return self.unit1.forward(x)
|
return self.unit1.forward(x)
|
||||||
@@ -90,9 +90,10 @@ if __name__ == "__main__":
|
|||||||
voc_size = 100
|
voc_size = 100
|
||||||
|
|
||||||
# Create Labeler
|
# Create Labeler
|
||||||
enc = Label(voc_size, 16, work_dir="../../results")
|
enc = Label(voc_size, 16, work_dir="../../results", encoding=Label.EncodingType.OneHot)
|
||||||
|
|
||||||
# Load
|
# Load
|
||||||
|
enc.fitter.init(0.1)
|
||||||
enc.fitter.load()
|
enc.fitter.load()
|
||||||
|
|
||||||
# Train
|
# Train
|
||||||
@@ -102,17 +103,18 @@ if __name__ == "__main__":
|
|||||||
# Save fitter state
|
# Save fitter state
|
||||||
enc.fitter.save()
|
enc.fitter.save()
|
||||||
|
|
||||||
|
# Test
|
||||||
# Encode test labels
|
# Encode test labels
|
||||||
test_labels = np.array([1, 23, 99, 37, 55, 7, 31, 10, 19, 70])
|
test_labels = train_labels
|
||||||
encoded, binary_labels = enc.encode(train_labels)
|
encoded, binary_labels = enc.encode(test_labels)
|
||||||
|
|
||||||
# Decode
|
# Decode
|
||||||
decoded_soft = enc.decode(encoded)
|
decoded_soft = enc.decode(encoded)
|
||||||
decode_hard = (decoded_soft > 0.9).astype(int)
|
decode_hard = (decoded_soft > 0.5).astype(int)
|
||||||
|
|
||||||
# Convert soft state into labels
|
# Convert soft state into labels
|
||||||
labels_reconst = enc.vec2label(decoded_soft)
|
labels_reconst = enc.vec2label(decoded_soft)
|
||||||
print(labels_reconst)
|
print(f"Reconst: {labels_reconst}")
|
||||||
|
|
||||||
print(f"Soft Error: {np.sum((decoded_soft - binary_labels) ** 2)}")
|
print(f"Soft Error: {np.sum((decoded_soft - binary_labels) ** 2)}")
|
||||||
print(f"Hard Error: {np.sum((decode_hard - binary_labels) ** 2)}")
|
print(f"Hard Error: {np.sum((decode_hard - binary_labels) ** 2)}")
|
||||||
|
|||||||
Reference in New Issue
Block a user