threshold-parity3 / create_safetensors.py
CharlesCNorton
3-bit parity threshold circuit, magnitude 20
93b5397
Raw
History Blame Contribute Delete
1.75 kB
import torch
from safetensors.torch import save_file
# Cascade of two XOR2 gates: parity(a,b,c) = XOR(XOR(a,b), c)
# Each XOR uses OR-NAND-AND structure
def xor_block(prefix):
return {
f'{prefix}.or.weight': torch.tensor([1.0, 1.0], dtype=torch.float32),
f'{prefix}.or.bias': torch.tensor([-1.0], dtype=torch.float32),
f'{prefix}.nand.weight': torch.tensor([-1.0, -1.0], dtype=torch.float32),
f'{prefix}.nand.bias': torch.tensor([1.0], dtype=torch.float32),
f'{prefix}.and.weight': torch.tensor([1.0, 1.0], dtype=torch.float32),
f'{prefix}.and.bias': torch.tensor([-2.0], dtype=torch.float32),
}
weights = {}
weights.update(xor_block('xor1'))
weights.update(xor_block('xor2'))
save_file(weights, 'model.safetensors')
def xor2(a, b, prefix):
or_out = int(a * weights[f'{prefix}.or.weight'][0] + b * weights[f'{prefix}.or.weight'][1] + weights[f'{prefix}.or.bias'] >= 0)
nand_out = int(a * weights[f'{prefix}.nand.weight'][0] + b * weights[f'{prefix}.nand.weight'][1] + weights[f'{prefix}.nand.bias'] >= 0)
and_out = int(or_out * weights[f'{prefix}.and.weight'][0] + nand_out * weights[f'{prefix}.and.weight'][1] + weights[f'{prefix}.and.bias'] >= 0)
return and_out
def parity3(a, b, c):
xor_ab = xor2(a, b, 'xor1')
return xor2(xor_ab, c, 'xor2')
print("Verifying parity3...")
errors = 0
for i in range(8):
a, b, c = (i >> 2) & 1, (i >> 1) & 1, i & 1
result = parity3(a, b, c)
expected = a ^ b ^ c
if result != expected:
errors += 1
print(f"ERROR: parity({a},{b},{c}) = {result}, expected {expected}")
if errors == 0:
print("All 8 test cases passed!")
print(f"Magnitude: {sum(t.abs().sum().item() for t in weights.values()):.0f}")