@@ -149,7 +149,10 @@ FrameLayout broadcast_layout(const std::vector<VALUETYPE>& values,
149149 const int nframes,
150150 const std::size_t stride,
151151 const char * name) {
152- if (values.empty () || values.size () == stride) {
152+ // A width of zero is a consumer that reads nothing. Whatever the caller
153+ // supplied was answered before the frames were reached, so they divide
154+ // nothing between them.
155+ if (stride == 0 || values.empty () || values.size () == stride) {
153156 return {stride, 0 };
154157 }
155158 return per_frame_layout (values, nframes, stride, name);
@@ -593,7 +596,8 @@ std::map<std::string, torch::Tensor> NativeSpinPTExpt::run_graph_payload(
593596 const std::int64_t nloc,
594597 const torch::Tensor& spin,
595598 const std::vector<double >& fparam,
596- const std::vector<double >& aparam) {
599+ const std::vector<double >& aparam,
600+ const std::vector<double >& charge_spin) {
597601 graph.edge_mask =
598602 deepmd::applyPairExclusion (graph.edge_index , graph.edge_mask , graph.atype ,
599603 pair_exclude_table_, ntypes);
@@ -607,13 +611,18 @@ std::map<std::string, torch::Tensor> NativeSpinPTExpt::run_graph_payload(
607611 deepmd::flatten_canonical_atom_virial (output_map);
608612 } else {
609613 const torch::Device device = graph.atype .device ();
614+ // This lower reads the condition as an ordinary input, so a call may name
615+ // the state it wants; the state the model was frozen with stands in for a
616+ // call that names none.
617+ const std::vector<double >& served =
618+ charge_spin.empty () ? default_chg_spin_ : charge_spin;
610619 extract_outputs (
611620 output_map,
612621 run_model_graph (
613622 graph, spin,
614623 make_fparam_tensor (fparam, default_fparam_, dfparam, device),
615624 make_aparam_tensor (aparam, daparam, node_count, nloc, device),
616- make_chg_spin_tensor (default_chg_spin_ , dchgspin, device)));
625+ make_chg_spin_tensor (served , dchgspin, device)));
617626 }
618627 return output_map;
619628}
@@ -651,6 +660,7 @@ void NativeSpinPTExpt::compute(ENERGYVTYPE& ener,
651660 const int & ago,
652661 const std::vector<VALUETYPE >& fparam,
653662 const std::vector<VALUETYPE >& aparam,
663+ const std::vector<double >& charge_spin,
654664 const bool atomic) {
655665 const torch::Device device = gpu_enabled ? torch::Device (torch::kCUDA , gpu_id)
656666 : torch::Device (torch::kCPU );
@@ -781,7 +791,7 @@ void NativeSpinPTExpt::compute(ENERGYVTYPE& ener,
781791 std::map<std::string, torch::Tensor> output_map = run_graph_payload (
782792 graph, node_count, nloc, spin_Tensor.slice (0 , 0 , node_count),
783793 std::vector<double >(fparam.begin (), fparam.end ()),
784- std::vector<double >(aparam_real.begin (), aparam_real.end ()));
794+ std::vector<double >(aparam_real.begin (), aparam_real.end ()), charge_spin );
785795
786796 // The forward emits flat per-node public keys; rewrite them into the dense
787797 // internal-key layout the extraction below reads. The extended node set
@@ -868,6 +878,7 @@ template void NativeSpinPTExpt::compute<double, std::vector<ENERGYTYPE>>(
868878 const int & ago,
869879 const std::vector<double >& fparam,
870880 const std::vector<double >& aparam,
881+ const std::vector<double >& charge_spin,
871882 const bool atomic);
872883template void NativeSpinPTExpt::compute<float , std::vector<ENERGYTYPE >>(
873884 std::vector<ENERGYTYPE >& ener,
@@ -884,6 +895,7 @@ template void NativeSpinPTExpt::compute<float, std::vector<ENERGYTYPE>>(
884895 const int & ago,
885896 const std::vector<float >& fparam,
886897 const std::vector<float >& aparam,
898+ const std::vector<double >& charge_spin,
887899 const bool atomic);
888900
889901// ============================================================================
@@ -903,6 +915,7 @@ void NativeSpinPTExpt::compute(ENERGYVTYPE& ener,
903915 const std::vector<VALUETYPE >& box,
904916 const std::vector<VALUETYPE >& fparam,
905917 const std::vector<VALUETYPE >& aparam,
918+ const std::vector<double >& charge_spin,
906919 const bool atomic) {
907920 const std::size_t nloc = atype.size ();
908921 const int nframes = frame_count (coord, nloc);
@@ -915,6 +928,8 @@ void NativeSpinPTExpt::compute(ENERGYVTYPE& ener,
915928 fparam, nframes, static_cast <std::size_t >(dfparam), " fparam" );
916929 const FrameLayout aparam_layout = broadcast_layout (
917930 aparam, nframes, nloc * static_cast <std::size_t >(daparam), " aparam" );
931+ const FrameLayout charge_spin_layout = broadcast_layout (
932+ charge_spin, nframes, static_cast <std::size_t >(dchgspin), " charge_spin" );
918933
919934 ener.clear ();
920935 force.clear ();
@@ -932,7 +947,8 @@ void NativeSpinPTExpt::compute(ENERGYVTYPE& ener,
932947 frame_block (spin, ff, spin_layout), atype,
933948 frame_block (box, ff, box_layout),
934949 frame_block (fparam, ff, fparam_layout),
935- frame_block (aparam, ff, aparam_layout), atomic);
950+ frame_block (aparam, ff, aparam_layout),
951+ frame_block (charge_spin, ff, charge_spin_layout), atomic);
936952 ener.insert (ener.end (), frame_ener.begin (), frame_ener.end ());
937953 force.insert (force.end (), frame_force.begin (), frame_force.end ());
938954 force_mag.insert (force_mag.end (), frame_force_mag.begin (),
@@ -958,6 +974,7 @@ void NativeSpinPTExpt::compute_frame(ENERGYVTYPE& ener,
958974 const std::vector<VALUETYPE >& box,
959975 const std::vector<VALUETYPE >& fparam,
960976 const std::vector<VALUETYPE >& aparam,
977+ const std::vector<double >& charge_spin,
961978 const bool atomic) {
962979 const torch::Device device = gpu_enabled ? torch::Device (torch::kCUDA , gpu_id)
963980 : torch::Device (torch::kCPU );
@@ -1036,10 +1053,10 @@ void NativeSpinPTExpt::compute_frame(ENERGYVTYPE& ener,
10361053 f64_options)
10371054 .clone ()
10381055 .to (device);
1039- std::map<std::string, torch::Tensor> output_map =
1040- run_graph_payload ( graph, nloc, nloc, spin_Tensor,
1041- std::vector<double >(fparam.begin (), fparam.end ()),
1042- std::vector<double >(aparam.begin (), aparam.end ()));
1056+ std::map<std::string, torch::Tensor> output_map = run_graph_payload (
1057+ graph, nloc, nloc, spin_Tensor,
1058+ std::vector<double >(fparam.begin (), fparam.end ()),
1059+ std::vector<double >(aparam.begin (), aparam.end ()), charge_spin );
10431060 deepmd::remap_graph_spin_outputs_to_dense_keys (output_map, nloc, nall,
10441061 atomic);
10451062
@@ -1110,6 +1127,7 @@ template void NativeSpinPTExpt::compute<double, std::vector<ENERGYTYPE>>(
11101127 const std::vector<double >& box,
11111128 const std::vector<double >& fparam,
11121129 const std::vector<double >& aparam,
1130+ const std::vector<double >& charge_spin,
11131131 const bool atomic);
11141132template void NativeSpinPTExpt::compute<float , std::vector<ENERGYTYPE >>(
11151133 std::vector<ENERGYTYPE >& ener,
@@ -1124,6 +1142,7 @@ template void NativeSpinPTExpt::compute<float, std::vector<ENERGYTYPE>>(
11241142 const std::vector<float >& box,
11251143 const std::vector<float >& fparam,
11261144 const std::vector<float >& aparam,
1145+ const std::vector<double >& charge_spin,
11271146 const bool atomic);
11281147
11291148// ============================================================================
@@ -1143,11 +1162,8 @@ void NativeSpinPTExpt::computew(std::vector<double>& ener,
11431162 const std::vector<double >& fparam,
11441163 const std::vector<double >& aparam,
11451164 const bool atomic) {
1146- translate_error ([&] {
1147- reject_unsupported_parametric_inputs (fparam, aparam, dfparam, daparam);
1148- compute (ener, force, force_mag, virial, atom_energy, atom_virial, coord,
1149- spin, atype, box, fparam, aparam, atomic);
1150- });
1165+ computew (ener, force, force_mag, virial, atom_energy, atom_virial, coord,
1166+ spin, atype, box, fparam, aparam, std::vector<double >(), atomic);
11511167}
11521168
11531169void NativeSpinPTExpt::computew (std::vector<double >& ener,
@@ -1163,11 +1179,8 @@ void NativeSpinPTExpt::computew(std::vector<double>& ener,
11631179 const std::vector<float >& fparam,
11641180 const std::vector<float >& aparam,
11651181 const bool atomic) {
1166- translate_error ([&] {
1167- reject_unsupported_parametric_inputs (fparam, aparam, dfparam, daparam);
1168- compute (ener, force, force_mag, virial, atom_energy, atom_virial, coord,
1169- spin, atype, box, fparam, aparam, atomic);
1170- });
1182+ computew (ener, force, force_mag, virial, atom_energy, atom_virial, coord,
1183+ spin, atype, box, fparam, aparam, std::vector<double >(), atomic);
11711184}
11721185
11731186void NativeSpinPTExpt::computew (std::vector<double >& ener,
@@ -1186,11 +1199,9 @@ void NativeSpinPTExpt::computew(std::vector<double>& ener,
11861199 const std::vector<double >& fparam,
11871200 const std::vector<double >& aparam,
11881201 const bool atomic) {
1189- translate_error ([&] {
1190- reject_unsupported_parametric_inputs (fparam, aparam, dfparam, daparam);
1191- compute (ener, force, force_mag, virial, atom_energy, atom_virial, coord,
1192- spin, atype, nghost, inlist, ago, fparam, aparam, atomic);
1193- });
1202+ computew (ener, force, force_mag, virial, atom_energy, atom_virial, coord,
1203+ spin, atype, box, nghost, inlist, ago, fparam, aparam,
1204+ std::vector<double >(), atomic);
11941205}
11951206
11961207void NativeSpinPTExpt::computew (std::vector<double >& ener,
@@ -1209,11 +1220,9 @@ void NativeSpinPTExpt::computew(std::vector<double>& ener,
12091220 const std::vector<float >& fparam,
12101221 const std::vector<float >& aparam,
12111222 const bool atomic) {
1212- translate_error ([&] {
1213- reject_unsupported_parametric_inputs (fparam, aparam, dfparam, daparam);
1214- compute (ener, force, force_mag, virial, atom_energy, atom_virial, coord,
1215- spin, atype, nghost, inlist, ago, fparam, aparam, atomic);
1216- });
1223+ computew (ener, force, force_mag, virial, atom_energy, atom_virial, coord,
1224+ spin, atype, box, nghost, inlist, ago, fparam, aparam,
1225+ std::vector<double >(), atomic);
12171226}
12181227
12191228void NativeSpinPTExpt::computew (std::vector<double >& ener,
@@ -1230,11 +1239,14 @@ void NativeSpinPTExpt::computew(std::vector<double>& ener,
12301239 const std::vector<double >& aparam,
12311240 const std::vector<double >& charge_spin,
12321241 const bool atomic) {
1233- check_call_charge_spin (
1234- charge_spin, frame_count (coord, atype.size ()), settable_chgspin,
1235- /* applied_per_call=*/ false , default_chg_spin_, chg_spin_table_ranges_);
1236- computew (ener, force, force_mag, virial, atom_energy, atom_virial, coord,
1237- spin, atype, box, fparam, aparam, atomic);
1242+ translate_error ([&] {
1243+ reject_unsupported_parametric_inputs (fparam, aparam, dfparam, daparam);
1244+ check_call_charge_spin (charge_spin, frame_count (coord, atype.size ()),
1245+ settable_chgspin, reads_charge_spin_per_call (),
1246+ default_chg_spin_, chg_spin_table_ranges_);
1247+ compute (ener, force, force_mag, virial, atom_energy, atom_virial, coord,
1248+ spin, atype, box, fparam, aparam, charge_spin, atomic);
1249+ });
12381250}
12391251
12401252void NativeSpinPTExpt::computew (std::vector<double >& ener,
@@ -1251,11 +1263,14 @@ void NativeSpinPTExpt::computew(std::vector<double>& ener,
12511263 const std::vector<float >& aparam,
12521264 const std::vector<double >& charge_spin,
12531265 const bool atomic) {
1254- check_call_charge_spin (
1255- charge_spin, frame_count (coord, atype.size ()), settable_chgspin,
1256- /* applied_per_call=*/ false , default_chg_spin_, chg_spin_table_ranges_);
1257- computew (ener, force, force_mag, virial, atom_energy, atom_virial, coord,
1258- spin, atype, box, fparam, aparam, atomic);
1266+ translate_error ([&] {
1267+ reject_unsupported_parametric_inputs (fparam, aparam, dfparam, daparam);
1268+ check_call_charge_spin (charge_spin, frame_count (coord, atype.size ()),
1269+ settable_chgspin, reads_charge_spin_per_call (),
1270+ default_chg_spin_, chg_spin_table_ranges_);
1271+ compute (ener, force, force_mag, virial, atom_energy, atom_virial, coord,
1272+ spin, atype, box, fparam, aparam, charge_spin, atomic);
1273+ });
12591274}
12601275
12611276void NativeSpinPTExpt::computew (std::vector<double >& ener,
@@ -1275,11 +1290,15 @@ void NativeSpinPTExpt::computew(std::vector<double>& ener,
12751290 const std::vector<double >& aparam,
12761291 const std::vector<double >& charge_spin,
12771292 const bool atomic) {
1278- check_call_charge_spin (charge_spin, 1 , settable_chgspin,
1279- /* applied_per_call=*/ false , default_chg_spin_,
1280- chg_spin_table_ranges_);
1281- computew (ener, force, force_mag, virial, atom_energy, atom_virial, coord,
1282- spin, atype, box, nghost, inlist, ago, fparam, aparam, atomic);
1293+ translate_error ([&] {
1294+ reject_unsupported_parametric_inputs (fparam, aparam, dfparam, daparam);
1295+ check_call_charge_spin (charge_spin, 1 , settable_chgspin,
1296+ reads_charge_spin_per_call (), default_chg_spin_,
1297+ chg_spin_table_ranges_);
1298+ compute (ener, force, force_mag, virial, atom_energy, atom_virial, coord,
1299+ spin, atype, nghost, inlist, ago, fparam, aparam, charge_spin,
1300+ atomic);
1301+ });
12831302}
12841303
12851304void NativeSpinPTExpt::computew (std::vector<double >& ener,
@@ -1299,11 +1318,15 @@ void NativeSpinPTExpt::computew(std::vector<double>& ener,
12991318 const std::vector<float >& aparam,
13001319 const std::vector<double >& charge_spin,
13011320 const bool atomic) {
1302- check_call_charge_spin (charge_spin, 1 , settable_chgspin,
1303- /* applied_per_call=*/ false , default_chg_spin_,
1304- chg_spin_table_ranges_);
1305- computew (ener, force, force_mag, virial, atom_energy, atom_virial, coord,
1306- spin, atype, box, nghost, inlist, ago, fparam, aparam, atomic);
1321+ translate_error ([&] {
1322+ reject_unsupported_parametric_inputs (fparam, aparam, dfparam, daparam);
1323+ check_call_charge_spin (charge_spin, 1 , settable_chgspin,
1324+ reads_charge_spin_per_call (), default_chg_spin_,
1325+ chg_spin_table_ranges_);
1326+ compute (ener, force, force_mag, virial, atom_energy, atom_virial, coord,
1327+ spin, atype, nghost, inlist, ago, fparam, aparam, charge_spin,
1328+ atomic);
1329+ });
13071330}
13081331
13091332void NativeSpinPTExpt::compute_canonical_graph_gpu (
0 commit comments