Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
470 changes: 470 additions & 0 deletions examples/11_hwa_lm_eval.ipynb

Large diffs are not rendered by default.

621 changes: 621 additions & 0 deletions examples/12_input_encoding_modes.ipynb

Large diffs are not rendered by default.

2 changes: 2 additions & 0 deletions pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -125,6 +125,8 @@ dependencies = [
"ninja>=1.11",
"qtorch>=0.3",
"smt>=2.9.4",
"transformers>=4.57.0",
"lm_eval>=0.4.9"
]

# List additional groups of dependencies here (e.g. development
Expand Down
4 changes: 2 additions & 2 deletions src/xbtorch/deployment/__init__.py
Original file line number Diff line number Diff line change
@@ -1,9 +1,9 @@
"""
Deployment (mapping, encoding, etc.) of solutions to inference accelerators
"""
from .base import Daffodil, SimpleFixedPoint
from .base import Daffodil, SimpleFixedPoint, register_accelerator
from .mapping import map_random
from .encoding import encode_simple_binary, encode_MAO, encode_LEA1, encode_LEA2
from .weight_encoding import encode_simple_binary, encode_MAO, encode_LEA1, encode_LEA2
from .metrics import compute_error
from .correction import train_collaborative, test_collaborative, CollaborativeLoss, add_collaborative_logistic_classifiers, dnn_favorable_searching_code
from .committee import test_committee
218 changes: 197 additions & 21 deletions src/xbtorch/deployment/base.py

Large diffs are not rendered by default.

2 changes: 1 addition & 1 deletion src/xbtorch/patches/__init__.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,4 @@
"""
Decorators for patching PyTorch models and optimizers for XBTorch
"""
from .model import xbtorch_model
from .model import xbtorch_model, replace_all_layers_stateless
221 changes: 180 additions & 41 deletions src/xbtorch/patches/decorators.py
Original file line number Diff line number Diff line change
Expand Up @@ -69,59 +69,198 @@ def backward_hook(module, grad_input, grad_output):
def xbtorch_forward(self, input):
# TODO: Add support for non-linear layers
if (self._xb_inference):
if (not hasattr(self, '_array_mappings')):
raise ValueError("Array mappings are not present, likely an issue during initialization.")
if (self.inference_accelerator.stateful):
if (not hasattr(self, '_array_mappings')):
raise ValueError("Array mappings are not present, likely an issue during initialization.")

gnorm_scale = 1.0
if (not self.inference_accelerator): raise ValueError('XB inference called without proper initialization of an accelerator profile.')

# ternary mapping scheme
v_read = self.inference_accelerator.v_read
g_norm = self.inference_accelerator.g_max - self.inference_accelerator.g_min

sw_weight = self.weight.data

gamma = torch.unique(sw_weight)[-1] # WAGE quantization learns matrices [-gamma, 0, gamma], and so it's important to scale either G matrices or input voltage vector

pos_idxs = self._array_mappings['Gpos']
neg_idxs = self._array_mappings['Gneg']

# pass to the crossbar, perform VMM, averaging/summing over multiple instances
pos_outputs = []
neg_outputs = []

if self.inference_accelerator.input_encoding_scheme == "instant":
# convert inputs to voltages, then quantize to DAC-based precision
input_voltages = input * v_read * gamma
input_voltages = self.inference_accelerator.DAC_quantize(input_voltages)

# Todo: can be possibly batched and made faster by using torch.bmm
for pos_idx in pos_idxs:
gpos = self.inference_accelerator.read_chip(pos_idx[0], sw_weight.shape[0], pos_idx[1], sw_weight.shape[1])
pos_outputs.append(input_voltages @ gpos.T)

for neg_idx in neg_idxs:
gneg = self.inference_accelerator.read_chip(neg_idx[0], sw_weight.shape[0], neg_idx[1], sw_weight.shape[1])
neg_outputs.append(input_voltages @ gneg.T)

# Convert the list of tensors to a single tensor
pos_outputs = torch.stack(pos_outputs)
neg_outputs = torch.stack(neg_outputs)

else:
# Linear temporal bit-sliced encoding
input_slices = self.inference_accelerator.DAC_quantize(
input * v_read * gamma
)

gnorm_scale = 1.0
if (not self.inference_accelerator): raise ValueError('XB inference called without proper initialization of an accelerator profile.')
pos_outputs_accum = []
neg_outputs_accum = []

for bit_idx, input_slice in enumerate(input_slices):
bit_weight = 2 ** bit_idx

pos_slice_outputs = []
neg_slice_outputs = []

for pos_idx in pos_idxs:
gpos = self.inference_accelerator.read_chip(
pos_idx[0],
sw_weight.shape[0],
pos_idx[1],
sw_weight.shape[1]
)
pos_slice_outputs.append(input_slice @ gpos.T)

for neg_idx in neg_idxs:
gneg = self.inference_accelerator.read_chip(
neg_idx[0],
sw_weight.shape[0],
neg_idx[1],
sw_weight.shape[1]
)
neg_slice_outputs.append(input_slice @ gneg.T)

pos_slice_outputs = torch.stack(pos_slice_outputs)
neg_slice_outputs = torch.stack(neg_slice_outputs)

pos_outputs_accum.append(bit_weight * pos_slice_outputs)
neg_outputs_accum.append(bit_weight * neg_slice_outputs)

pos_outputs = torch.stack(pos_outputs_accum).sum(dim=0)
neg_outputs = torch.stack(neg_outputs_accum).sum(dim=0)

# normalize temporal accumulation
scale = (2 ** self.inference_accelerator.dac_bits - 1)
pos_outputs = pos_outputs / scale
neg_outputs = neg_outputs / scale

# for MAO, this has to be sum, for regular mapping, this will be average
if (self._array_mappings['output_polling_mode'] == 'avg'):
output = torch.mean(pos_outputs, dim=0) - torch.mean(neg_outputs, dim=0)
elif (self._array_mappings['output_polling_mode'] == 'sum'):
output = torch.sum(pos_outputs, dim=0) - torch.sum(neg_outputs, dim=0)
elif (self._array_mappings['output_polling_mode'] == 'reduced_avg'):
output = torch.sum(pos_outputs * self._array_mappings['maskpos'], dim=0) / self._array_mappings['alpha'] - torch.sum(neg_outputs * self._array_mappings['maskneg'], dim=0) / self._array_mappings['alpha']
else:
raise ValueError("output_polling_mode not implemented")
# readback currents from ADC by simulated quantization again
output = self.inference_accelerator.ADC_quantize(output) # equivalent to optimizing the TIA potentiometer resistance. ADC quantization
output = output / (gnorm_scale * g_norm * v_read)

if (self.bias is not None): output += self.bias.data
return output
else:
# stateless forward, weights were never encoded and mapped to a crossbar. Here, we'll do everything on-the-fly.

# ternary mapping scheme
v_read = self.inference_accelerator.v_read
g_norm = self.inference_accelerator.g_max - self.inference_accelerator.g_min
gnorm_scale = 1.0
if (not self.inference_accelerator): raise ValueError('XB inference called without proper initialization of an accelerator profile.')

sw_weight = self.weight.data
# ternary mapping scheme
v_read = self.inference_accelerator.v_read
g_norm = self.inference_accelerator.g_max - self.inference_accelerator.g_min

gamma = torch.unique(sw_weight)[-1] # WAGE quantization learns matrices [-gamma, 0, gamma], and so it's important to scale either G matrices or input voltage vector
sw_weight = self.weight.data

# convert inputs to voltages, then quantize to DAC-based precision
input_voltages = input * v_read * gamma
input_voltages = self.inference_accelerator.DAC_quantize(input_voltages)
gamma = torch.unique(sw_weight)[-1] # WAGE quantization learns matrices [-gamma, 0, gamma], and so it's important to scale either G matrices or input voltage vector

pos_idxs = self._array_mappings['Gpos']
neg_idxs = self._array_mappings['Gneg']

# pass to the crossbar, perform VMM, averaging/summing over multiple instances
pos_outputs = []
neg_outputs = []

# Todo: can be possibly batched and made faster by using torch.bmm
for pos_idx in pos_idxs:
gpos = self.inference_accelerator.read_chip(pos_idx[0], sw_weight.shape[0], pos_idx[1], sw_weight.shape[1])
pos_outputs.append(input_voltages @ gpos.T)
# stateless operation
# first, let's convert weights to conductances
# TODO: raise error if redundnancy based scheme required with stateless operation, not supported atm
# would require splitting initialize layer mappings for stateless and stateful

Gposs, Gnegs = self.inference_accelerator.map_weights_to_array_stateless(sw_weight)

# pass to the crossbar, perform VMM, averaging/summing over multiple instances
pos_outputs = []
neg_outputs = []

if self.inference_accelerator.input_encoding_scheme == "instant":
# convert inputs to voltages, then quantize to DAC-based precision
input_voltages = input * v_read * gamma
input_voltages = self.inference_accelerator.DAC_quantize(input_voltages)

# Todo: can be possibly batched and made faster by using torch.bmm
for Gpos in Gposs:
gpos = self.inference_accelerator.read_chip_stateless(Gpos)
pos_outputs.append(input_voltages @ gpos.T)

for Gneg in Gnegs:
gneg = self.inference_accelerator.read_chip_stateless(Gneg)
neg_outputs.append(input_voltages @ gneg.T)

# Convert the list of tensors to a single tensor
pos_outputs = torch.stack(pos_outputs)
neg_outputs = torch.stack(neg_outputs)

else:
# Linear temporal bit-sliced encoding
input_slices = self.inference_accelerator.DAC_quantize(
input * v_read * gamma
)

pos_outputs_accum = []
neg_outputs_accum = []

for neg_idx in neg_idxs:
gneg = self.inference_accelerator.read_chip(neg_idx[0], sw_weight.shape[0], neg_idx[1], sw_weight.shape[1])
neg_outputs.append(input_voltages @ gneg.T)
for bit_idx, input_slice in enumerate(input_slices):
bit_weight = 2 ** bit_idx

# Convert the list of tensors to a single tensor
pos_outputs = torch.stack(pos_outputs)
neg_outputs = torch.stack(neg_outputs)
pos_slice_outputs = []
neg_slice_outputs = []

for Gpos in Gposs:
gpos = self.inference_accelerator.read_chip_stateless(Gpos)
pos_slice_outputs.append(input_slice @ gpos.T)

for Gneg in Gnegs:
gneg = self.inference_accelerator.read_chip_stateless(Gneg)
neg_slice_outputs.append(input_slice @ gneg.T)

pos_slice_outputs = torch.stack(pos_slice_outputs)
neg_slice_outputs = torch.stack(neg_slice_outputs)

pos_outputs_accum.append(bit_weight * pos_slice_outputs)
neg_outputs_accum.append(bit_weight * neg_slice_outputs)

pos_outputs = torch.stack(pos_outputs_accum).sum(dim=0)
neg_outputs = torch.stack(neg_outputs_accum).sum(dim=0)

# normalize temporal accumulation
scale = (2 ** self.inference_accelerator.dac_bits - 1)
pos_outputs = pos_outputs / scale
neg_outputs = neg_outputs / scale

# for MAO, this has to be sum, for regular mapping, this will be average
if (self._array_mappings['output_polling_mode'] == 'avg'):
output = torch.mean(pos_outputs, dim=0) - torch.mean(neg_outputs, dim=0)
elif (self._array_mappings['output_polling_mode'] == 'sum'):
output = torch.sum(pos_outputs, dim=0) - torch.sum(neg_outputs, dim=0)
elif (self._array_mappings['output_polling_mode'] == 'reduced_avg'):
output = torch.sum(pos_outputs * self._array_mappings['maskpos'], dim=0) / self._array_mappings['alpha'] - torch.sum(neg_outputs * self._array_mappings['maskneg'], dim=0) / self._array_mappings['alpha']
else:
raise ValueError("output_polling_mode not implemented")
# readback currents from ADC by simulated quantization again
output = self.inference_accelerator.ADC_quantize(output) # equivalent to optimizing the TIA potentiometer resistance. ADC quantization
output = output / (gnorm_scale * g_norm * v_read)
# TODO: output polling modes are not implemented for now

if (self.bias is not None): output += self.bias.data
return output
# readback currents from ADC by simulated quantization again
output = self.inference_accelerator.ADC_quantize(output) # equivalent to optimizing the TIA potentiometer resistance. ADC quantization
output = output / (gnorm_scale * g_norm * v_read)

if (self.bias is not None): output += self.bias.data
return output
else:
if (hasattr(self, 'weight')):
self.weight.input = input
Expand Down
Loading
Loading