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
2 changes: 1 addition & 1 deletion cpp/include/seqwin/filter.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -12,7 +12,7 @@ void get_penalty(
const Kmer* kmers,
Node* nodes,
std::size_t n_nodes,
const std::size_t* record_offsets,
const std::uint32_t* record_offsets,
std::size_t n_record_offsets,
const bool* is_targets,
std::size_t n_assemblies,
Expand Down
2 changes: 1 addition & 1 deletion cpp/include/seqwin/graph.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -66,7 +66,7 @@ struct Graph {
/** Sorted by hash. */
NoInitArray<Edge> edges;
/** Cumulative global FASTA record offsets by assembly. */
std::vector<std::size_t> record_offsets;
std::vector<std::uint32_t> record_offsets;
/** FASTA record IDs of each assembly. */
std::vector<std::vector<std::string>> ids_by_assembly;
};
Expand Down
4 changes: 2 additions & 2 deletions cpp/src/bindings/python_bindings.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -92,7 +92,7 @@ PYBIND11_MODULE(_core, m) {
m.def("_get_penalty_native",
[](py::array_t<seqwin::Kmer, py::array::c_style> kmers,
py::array_t<seqwin::Node, py::array::c_style> nodes,
py::array_t<std::size_t, py::array::c_style> record_offsets,
py::array_t<std::uint32_t, py::array::c_style> record_offsets,
py::array_t<bool, py::array::c_style> is_targets,
std::size_t n_cpu
) {
Expand All @@ -107,7 +107,7 @@ PYBIND11_MODULE(_core, m) {

const auto* kmers_ptr = static_cast<const seqwin::Kmer*>(kmers_buf.ptr);
auto* nodes_ptr = static_cast<seqwin::Node*>(nodes_buf.ptr);
const auto* record_offsets_ptr = static_cast<const std::size_t*>(record_offsets_buf.ptr);
const auto* record_offsets_ptr = static_cast<const std::uint32_t*>(record_offsets_buf.ptr);
const auto* is_targets_ptr = static_cast<const bool*>(is_targets_buf.ptr);
const auto n_nodes = static_cast<std::size_t>(nodes_buf.shape[0]);
const auto n_record_offsets = static_cast<std::size_t>(record_offsets_buf.shape[0]);
Expand Down
19 changes: 13 additions & 6 deletions cpp/src/seqwin/build.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -132,8 +132,13 @@ ThreadGraph build_worker(
std::vector<std::string> record_ids;
record_ids.reserve(records.size());

std::uint32_t record_idx = graph.record_offsets.back();
if (records.size() > std::numeric_limits<std::uint32_t>::max() - record_idx) {
throw std::runtime_error(
"Total number of FASTA records exceeds uint32 range"
);
}
for (std::size_t record_i = 0; record_i < records.size(); ++record_i) {
const auto record_idx = graph.record_offsets.back() + record_i;
auto& record = records[record_i];
if (record.sequence.size() > std::numeric_limits<std::uint32_t>::max()) {
throw std::runtime_error(
Expand All @@ -151,7 +156,7 @@ ThreadGraph build_worker(
m.out_hash,
Kmer{
static_cast<std::uint32_t>(m.pos),
static_cast<std::uint32_t>(record_idx)
record_idx
}
});
}
Expand All @@ -161,6 +166,7 @@ ThreadGraph build_worker(
++node_it->second;
++graph.n_kmers;
}
++record_idx;

// Current record is too short for an edge
if (mins.size() < 2) {
Expand All @@ -182,7 +188,7 @@ ThreadGraph build_worker(
}
}
}
graph.record_offsets.push_back(graph.record_offsets.back() + records.size());
graph.record_offsets.push_back(record_idx);
graph.ids_by_assembly.push_back(std::move(record_ids));
}

Expand Down Expand Up @@ -260,7 +266,7 @@ NoInitArray<Kmer> recompute_kmers(
std::size_t kmerlen,
std::size_t windowsize,
const std::vector<ThreadGraph>& graphs,
const std::vector<std::size_t>& record_offsets,
const std::vector<std::uint32_t>& record_offsets,
KmerMaps& kmer_maps,
ThreadPool& pool
) {
Expand All @@ -279,9 +285,9 @@ NoInitArray<Kmer> recompute_kmers(
assembly_i < graph.end_assembly;
++assembly_i) {
auto records = read_fasta(assembly_paths[assembly_i]);
std::uint32_t record_idx = record_offsets[assembly_i];

for (std::size_t record_i = 0; record_i < records.size(); ++record_i) {
const auto record_idx = record_offsets[assembly_i] + record_i;
auto& record = records[record_i];
if (record.sequence.size() > std::numeric_limits<std::uint32_t>::max()) {
throw std::runtime_error(
Expand All @@ -305,9 +311,10 @@ NoInitArray<Kmer> recompute_kmers(
auto& cursor = cursor_it->second;
kmers[cursor++] = Kmer{
static_cast<std::uint32_t>(m.pos),
static_cast<std::uint32_t>(record_idx)
record_idx
};
}
++record_idx;
}
}
KmerMap{}.swap(hash_to_cursor);
Expand Down
18 changes: 7 additions & 11 deletions cpp/src/seqwin/build_internals.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -300,10 +300,6 @@ std::pair<Graph, KmerMaps> merge_thread_graphs(
) {
if (graphs.size() == 1) {
auto& graph = graphs[0];
if (graph.record_offsets.back() > std::numeric_limits<std::uint32_t>::max()) {
throw std::runtime_error("Total number of FASTA records exceeds uint32 range");
}

merge_edges(graph.edges, pool); // Sort only

auto merged = merge_nodes(graph.nodes, graphs, pool, low_memory);
Expand Down Expand Up @@ -337,25 +333,25 @@ std::pair<Graph, KmerMaps> merge_thread_graphs(

// Merge record offsets
std::vector<std::uint32_t> thread_record_offsets(graphs.size());
std::vector<std::size_t> record_offsets;
std::vector<std::uint32_t> record_offsets;
record_offsets.reserve(n_assemblies + 1);
record_offsets.push_back(0);

std::size_t total_records = 0;
std::uint32_t total_records = 0;
for (std::size_t t = 0; t < graphs.size(); ++t) {
auto& local_offsets = graphs[t].record_offsets;

const auto base = total_records;
thread_record_offsets[t] = static_cast<std::uint32_t>(base);
thread_record_offsets[t] = base;
if (local_offsets.back() > std::numeric_limits<std::uint32_t>::max() - total_records) {
throw std::runtime_error("Total number of FASTA records exceeds uint32 range");
}
total_records += local_offsets.back();

for (std::size_t i = 1; i < local_offsets.size(); ++i) {
record_offsets.push_back(base + local_offsets[i]);
}
std::vector<std::size_t>().swap(local_offsets);
}
if (total_records > std::numeric_limits<std::uint32_t>::max()) {
throw std::runtime_error("Total number of FASTA records exceeds uint32 range");
std::vector<std::uint32_t>().swap(local_offsets);
}

// Merge edges and nodes first to reduce peak memory
Expand Down
2 changes: 1 addition & 1 deletion cpp/src/seqwin/build_internals.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -44,7 +44,7 @@ struct ThreadGraph {
/** Unsorted. */
NoInitArray<Edge> edges;
/** Thread-local cumulative FASTA record offsets by assembly. */
std::vector<std::size_t> record_offsets;
std::vector<std::uint32_t> record_offsets;
/** FASTA record IDs of each assembly. */
std::vector<std::vector<std::string>> ids_by_assembly;
/** Total number of minimizers generated by this worker. */
Expand Down
97 changes: 55 additions & 42 deletions cpp/src/seqwin/filter.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -16,12 +16,20 @@ void get_penalty(
const Kmer* kmers,
Node* nodes,
std::size_t n_nodes,
const std::size_t* record_offsets,
const std::uint32_t* record_offsets,
std::size_t n_record_offsets,
const bool* is_targets,
std::size_t n_assemblies,
std::size_t n_cpu
) {
/** Metadata shared by all FASTA records in one assembly. */
struct RecordInfo {
/** Inclusive global index of the assembly's final FASTA record. */
std::uint32_t last_record_idx;
/** Whether the assembly belongs to the target set, as 0 or 1. */
std::uint32_t is_target;
};

if (n_record_offsets != n_assemblies + 1) {
throw std::invalid_argument("len(record_offsets) must equal len(is_targets) + 1");
}
Expand Down Expand Up @@ -57,6 +65,30 @@ void get_penalty(
}
internal::ThreadPool pool(n_workers);

const std::uint32_t n_records = record_offsets[n_assemblies];
NoInitArray<RecordInfo> record_info(n_records);
pool.parallel_for(n_assemblies, [&](std::size_t start, std::size_t end, std::size_t) {
for (std::size_t assembly_idx = start; assembly_idx < end; ++assembly_idx) {
const std::uint32_t record_start = record_offsets[assembly_idx];
const std::uint32_t record_stop = record_offsets[assembly_idx + 1];
if (record_start == record_stop) {
continue;
}
const RecordInfo info{
record_stop - 1,
is_targets[assembly_idx] ? 1U : 0U
};
std::fill(
record_info.begin() + record_start,
record_info.begin() + record_stop,
info
);
}
});

const double inv_total_targets = 1.0 / static_cast<double>(total_targets);
const double inv_total_non_targets = 1.0 / static_cast<double>(total_non_targets);

pool.parallel_for(n_nodes, [&](std::size_t start, std::size_t end, std::size_t) {
for (std::size_t node_i = start; node_i < end; ++node_i) {
auto& node = nodes[node_i];
Expand All @@ -67,58 +99,39 @@ void get_penalty(
continue;
}

std::uint32_t n_tar = 0;
std::uint32_t n_neg = 0;

// Monotonic scan of record_idx and record_offsets
// Each node range has nondecreasing record_idx values, so one upper_bound()
// maps the first record to its assembly and the scan only advances forward
std::uint32_t previous_record_idx = kmers[node.start].record_idx;
if (static_cast<std::size_t>(previous_record_idx) >= record_offsets[n_assemblies]) {
auto previous_record_idx = kmers[node.start].record_idx;
if (previous_record_idx >= n_records) {
throw std::invalid_argument("record_idx is outside record_offsets range");
}
const auto* offset_it = std::upper_bound(
record_offsets,
record_offsets + n_record_offsets,
static_cast<std::size_t>(previous_record_idx)
);
std::size_t assembly_idx = static_cast<std::size_t>(offset_it - record_offsets - 1);
// Count each assembly once per node
std::size_t last_counted_assembly = std::numeric_limits<std::size_t>::max();

for (std::size_t kmer_i = node.start; kmer_i < node.stop; ++kmer_i) {
const std::uint32_t record_idx_u32 = kmers[kmer_i].record_idx;
if (record_idx_u32 < previous_record_idx) {
auto info = record_info[previous_record_idx];
auto last_record_idx = info.last_record_idx;
std::uint32_t n_tar = info.is_target;
std::uint32_t n_neg = 1U - info.is_target;

for (std::size_t kmer_i = node.start + 1; kmer_i < node.stop; ++kmer_i) {
const std::uint32_t record_idx = kmers[kmer_i].record_idx;
if (record_idx < previous_record_idx) {
throw std::invalid_argument("record_idx must be nondecreasing within each node range");
}
previous_record_idx = record_idx_u32;
previous_record_idx = record_idx;

const std::size_t record_idx = static_cast<std::size_t>(record_idx_u32);
if (record_idx >= record_offsets[n_assemblies]) {
throw std::invalid_argument("record_idx is outside record_offsets range");
}
// Duplicate record offsets are allowed for zero-record assemblies
while (
assembly_idx + 1 < n_assemblies &&
record_idx >= record_offsets[assembly_idx + 1]
) {
++assembly_idx;
if (record_idx <= last_record_idx) {
continue;
}
if (assembly_idx != last_counted_assembly) {
if (is_targets[assembly_idx]) {
++n_tar;
} else {
++n_neg;
}
last_counted_assembly = assembly_idx;
if (record_idx >= n_records) {
throw std::invalid_argument("record_idx is outside record_offsets range");
}
info = record_info[record_idx];
last_record_idx = info.last_record_idx;
n_tar += info.is_target;
n_neg += 1U - info.is_target;
}

node.n_tar = n_tar;
node.n_neg = n_neg;
const double frac_tar = static_cast<double>(n_tar) / static_cast<double>(total_targets);
const double frac_neg = static_cast<double>(n_neg) / static_cast<double>(total_non_targets);
node.penalty = std::hypot(1.0 - frac_tar, frac_neg);
const double frac_tar = static_cast<double>(n_tar) * inv_total_targets;
const double frac_neg = static_cast<double>(n_neg) * inv_total_non_targets;
node.penalty = std::sqrt((1.0 - frac_tar) * (1.0 - frac_tar) + frac_neg * frac_neg);
}
});
}
Expand Down
6 changes: 3 additions & 3 deletions src/seqwin/graph/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -130,7 +130,7 @@ def build(
- 'first' (uint64): Smaller endpoint hash of the undirected edge.
- 'second' (uint64): Larger endpoint hash of the undirected edge.
- 'weight' (uintp): Number of assemblies where the endpoints are adjacent.
4. NDArray[np.uintp]: Cumulative global FASTA record offsets by assembly.
4. NDArray[np.uint32]: Cumulative global FASTA record offsets by assembly.
5. list[tuple[str, ...]]: FASTA record IDs of each assembly.
"""
return _build_native(
Expand All @@ -145,7 +145,7 @@ def build(
def _get_penalty(
kmers: NDArray[np.void],
nodes: NDArray[np.void],
record_offsets: NDArray[np.uintp],
record_offsets: NDArray[np.uint32],
is_targets: Iterable[bool],
n_cpu: int = 1
) -> None:
Expand All @@ -154,7 +154,7 @@ def _get_penalty(
Args:
kmers (NDArray): See `KmerGraph.kmers`.
nodes (NDArray): See `KmerGraph.nodes`.
record_offsets (NDArray[np.uintp]): See `KmerGraph.record_offsets`.
record_offsets (NDArray[np.uint32]): See `KmerGraph.record_offsets`.
is_targets (Iterable[bool]): Whether each assembly is a target assembly.
n_cpu (int, optional): Number of worker threads to use. [1]
"""
Expand Down
4 changes: 2 additions & 2 deletions src/seqwin/kmers.py
Original file line number Diff line number Diff line change
Expand Up @@ -54,7 +54,7 @@ class KmerGraph(object):
For each node, `kmers[node['start']:node['stop']]` is the k-mer group with `node['hash']`.
edges (NDArray[np.void]): A 1-D NumPy structured array of weighted, undirected edges.
Edge weight is the number of assemblies where the two k-mers are adjacent.
record_offsets (NDArray[np.uintp]): Cumulative global FASTA record offsets by assembly.
record_offsets (NDArray[np.uint32]): Cumulative global FASTA record offsets by assembly.
graph (nx.Graph): The graph instance built from filtered nodes and edges.
subgraphs (tuple[frozenset[np.uint64], ...] | None): Low-penalty subgraphs. Each subgraph is a set of k-mer hash values.
Generated with `self.filter()`.
Expand All @@ -65,7 +65,7 @@ class KmerGraph(object):
kmers: NDArray[np.void]
nodes: NDArray[np.void]
edges: NDArray[np.void]
record_offsets: NDArray[np.uintp]
record_offsets: NDArray[np.uint32]
graph: nx.Graph
subgraphs: tuple[frozenset[np.uint64], ...] | None
_is_filtered: bool # True if `self.filter()` is called
Expand Down
2 changes: 1 addition & 1 deletion src/seqwin/markers.py
Original file line number Diff line number Diff line change
Expand Up @@ -357,7 +357,7 @@ def _create_ck(
graph: nx.Graph,
nodes: tuple[np.uint64],
kmers: tuple,
record_offsets: NDArray[np.uintp],
record_offsets: NDArray[np.uint32],
n_tar: int,
kmerlen: int,
windowsize: int
Expand Down
Binary file modified tests/smoke/fixtures/expected/graph.npz
Binary file not shown.
Loading