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
28 changes: 28 additions & 0 deletions qmp/models/hubbard.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,7 @@
from ..networks.mlp import WaveFunctionElectron as MlpWaveFunctionElectron
from ..networks.transformers import WaveFunctionElectronUpDown as TransformersWaveFunction
from ..networks.transformers import WaveFunctionElectron as TransformersWaveFunctionElectron
from ..networks.mps import WaveFunction as MpsWaveFunction
from ..hamiltonian import Hamiltonian
from ..utility.model_dict import model_dict, ModelProto, NetworkProto, NetworkConfigProto

Expand Down Expand Up @@ -359,3 +360,30 @@ def create(self, model: Model) -> NetworkProto:


Model.network_dict["transformers/u1"] = TransformersElectronConfig


@dataclasses.dataclass
class MpsConfig:
"""
The configuration of the MPS network.
"""

# The bond dimension of the MPS
virtual_dim: int = 16

def create(self, model: Model) -> NetworkProto:
"""
Create an MPS network for the model.
"""
logging.info("MPS bond dimension: %d", self.virtual_dim)

network = MpsWaveFunction(
sites=model.m * model.n * 2,
physical_dim=2,
virtual_dim=self.virtual_dim,
)

return network


Model.network_dict["mps"] = MpsConfig
28 changes: 28 additions & 0 deletions qmp/models/ising.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,7 @@
from ..networks.mlp import WaveFunctionElectron as MlpWaveFunctionElectron
from ..networks.transformers import WaveFunctionNormal as TransformersWaveFunction
from ..networks.transformers import WaveFunctionElectron as TransformersWaveFunctionElectron
from ..networks.mps import WaveFunction as MpsWaveFunction
from ..hamiltonian import Hamiltonian
from ..utility.model_dict import model_dict, ModelProto, NetworkProto, NetworkConfigProto

Expand Down Expand Up @@ -408,3 +409,30 @@ def create(self, model: Model) -> NetworkProto:


Model.network_dict["transformers/u1"] = TransformersElectronConfig


@dataclasses.dataclass
class MpsConfig:
"""
The configuration of the MPS network.
"""

# The bond dimension of the MPS
virtual_dim: int = 16

def create(self, model: Model) -> NetworkProto:
"""
Create an MPS network for the model.
"""
logging.info("MPS bond dimension: %d", self.virtual_dim)

network = MpsWaveFunction(
sites=model.m * model.n,
physical_dim=2,
virtual_dim=self.virtual_dim,
)

return network


Model.network_dict["mps"] = MpsConfig
Loading