Skip to content

Commit cd3fa56

Browse files
committed
fix(api_cc): let a native-spin call name the charge state it wants served
The two artifacts this backend serves carry the condition differently, and the charge-aware entry point treated both as the folded one: it refused any state but the installed default and then dropped the argument, leaving run_graph_payload to build the input from that default. For a compressed artifact this is the contract. Its condition lives in frozen tables that set_charge_spin rebuilds and that stand for the whole run, and its compact lower takes no condition at all. Nothing there changes. An uncompressed one keeps the condition in the argument list of its compiled forward, which was already being filled on every evaluation, just from the default rather than from the caller. The condition now travels with the frames it belongs to, divided by the layout the other per-frame inputs use, so successive calls may each name their own state. The width of the forward's argument is what tells the two apart, and the charge-aware overload becomes the implementation the charge-unaware one delegates to, so the condition is an ordinary input rather than a second entry point.
1 parent e3480cf commit cd3fa56

3 files changed

Lines changed: 124 additions & 68 deletions

File tree

source/api_cc/include/NativeSpinPTExpt.h

Lines changed: 17 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -347,6 +347,7 @@ class NativeSpinPTExpt : public DeepSpinBackend {
347347
const int& ago,
348348
const std::vector<VALUETYPE>& fparam,
349349
const std::vector<VALUETYPE>& aparam,
350+
const std::vector<double>& charge_spin,
350351
const bool atomic);
351352
/**
352353
* @brief Evaluate frames that arrive without a neighbor list.
@@ -369,8 +370,21 @@ class NativeSpinPTExpt : public DeepSpinBackend {
369370
const std::vector<VALUETYPE>& box,
370371
const std::vector<VALUETYPE>& fparam,
371372
const std::vector<VALUETYPE>& aparam,
373+
const std::vector<double>& charge_spin,
372374
const bool atomic);
373375

376+
/**
377+
* @brief Whether a call may name the charge state it wants served.
378+
*
379+
* The two artifacts this class serves carry the condition differently. An
380+
* uncompressed one keeps it in the argument list of its compiled forward,
381+
* where a state reaches the evaluation that names it. A compressed one
382+
* folds it into frozen tables, which are rebuilt by ``set_charge_spin``
383+
* and stand for the whole run, so a call can only restate what they hold.
384+
* The width of the forward's argument tells the two apart.
385+
**/
386+
bool reads_charge_spin_per_call() const { return dchgspin > 0; }
387+
374388
/**
375389
* @brief Evaluate one frame: build a neighbor list, then fold ghost
376390
* contributions back onto their local owners.
@@ -388,6 +402,7 @@ class NativeSpinPTExpt : public DeepSpinBackend {
388402
const std::vector<VALUETYPE>& box,
389403
const std::vector<VALUETYPE>& fparam,
390404
const std::vector<VALUETYPE>& aparam,
405+
const std::vector<double>& charge_spin,
391406
const bool atomic);
392407

393408
/**
@@ -444,7 +459,8 @@ class NativeSpinPTExpt : public DeepSpinBackend {
444459
const std::int64_t nloc,
445460
const torch::Tensor& spin,
446461
const std::vector<double>& fparam,
447-
const std::vector<double>& aparam);
462+
const std::vector<double>& aparam,
463+
const std::vector<double>& charge_spin);
448464

449465
/**
450466
* @brief Bind the flat artifact outputs to their metadata key names.

source/api_cc/src/NativeSpinPTExpt.cc

Lines changed: 72 additions & 49 deletions
Original file line numberDiff line numberDiff line change
@@ -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);
872883
template 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);
11141132
template 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

11531169
void 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

11731186
void 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

11961207
void 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

12191228
void 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

12401252
void 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

12611276
void 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

12851304
void 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

13091332
void NativeSpinPTExpt::compute_canonical_graph_gpu(

0 commit comments

Comments
 (0)