Convolution configuration

class torch_lattice.nn.functional.conv.conv_config.ConvConfig(dataflow=Dataflow.ImplicitGEMM, ifsort=False, kmap_mode='hashmap_on_the_fly', downsample_mode='spconv', split_mask_num=1, split_mask_num_bwd=3, wgrad_split_k='auto', IGEMM_center_only=False, epsilon=0.0, mm_thresh=0, FOD_fusion=False)[source]

Bases: Mapping[str, Any]

Validated sparse-convolution execution policy.

Parameters:
dataflow: Dataflow
ifsort: bool
kmap_mode: str
downsample_mode: str
split_mask_num: int
split_mask_num_bwd: int
wgrad_split_k: int | str
IGEMM_center_only: bool
epsilon: float
mm_thresh: int
FOD_fusion: bool
copy()[source]
Return type:

ConvConfig

class torch_lattice.nn.functional.conv.conv_config.Dataflow(*values)[source]

Bases: Enum

ImplicitGEMM = 0
GatherScatter = 1
FetchOnDemand = 2
torch_lattice.nn.functional.conv.conv_config.clear_global_conv_config()[source]
Return type:

None

torch_lattice.nn.functional.conv.conv_config.get_default_conv_config(conv_mode=ConvMode.mode0, training=False)[source]
Return type:

ConvConfig

Parameters:
torch_lattice.nn.functional.conv.conv_config.get_global_conv_config()[source]
Return type:

ConvConfig | None

torch_lattice.nn.functional.conv.conv_config.set_global_conv_config(conv_config)[source]
Return type:

None

Parameters:

conv_config (ConvConfig | Mapping[str, Any])

class torch_lattice.nn.functional.conv.conv_mode.ConvMode(*values)[source]

Bases: Enum

mode0 = 0
mode1 = 1
mode2 = 2
torch_lattice.nn.functional.conv.conv_mode.get_conv_mode()[source]
Return type:

ConvMode

torch_lattice.nn.functional.conv.conv_mode.set_conv_mode(conv_mode)[source]
Return type:

None

Parameters:

conv_mode (int | ConvMode)