diff --git a/qmp/models/hubbard.py b/qmp/models/hubbard.py index 89526c6..c89828c 100644 --- a/qmp/models/hubbard.py +++ b/qmp/models/hubbard.py @@ -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 @@ -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 diff --git a/qmp/models/ising.py b/qmp/models/ising.py index 6a2a4fd..be2f636 100644 --- a/qmp/models/ising.py +++ b/qmp/models/ising.py @@ -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 @@ -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