diff --git a/lib/ancestor_builder.c b/lib/ancestor_builder.c index 2d36232e..bfbf920d 100644 --- a/lib/ancestor_builder.c +++ b/lib/ancestor_builder.c @@ -467,18 +467,6 @@ ancestor_builder_compute_ancestral_states(const ancestor_builder_t *self, int di } else { consensus = 0; } - /* printf("\t:ones=%d, consensus=%d\n", (int) ones, consensus); */ - /* fflush(stdout); */ - for (j = 0; j < sample_set_size; j++) { - u = sample_set[j]; - if (disagree[u] && (genotypes[j] != consensus) - && (genotypes[j] != TSK_MISSING_DATA)) { - /* This sample has disagreed with consensus twice in a row, - * so remove it */ - /* printf("\t\tremoving %d\n", sample_set[j]); */ - sample_set[j] = -1; - } - } site_time = sites[l].time; if (site_time > focal_site_time) { @@ -488,31 +476,47 @@ ancestor_builder_compute_ancestral_states(const ancestor_builder_t *self, int di ancestor[l] = consensus; } } - /* For the remaining samples, set the disagree flags based - * on whether they agree with the consensus for this site. */ + derived_count = sites[l].derived_count; if ((site_time > focal_site_time) || (derived_count > ones)) { - for (j = 0; j < sample_set_size; j++) { - u = sample_set[j]; - if (u != -1) { - disagree[u] = ((genotypes[j] != consensus) - && (genotypes[j] != TSK_MISSING_DATA)); + if (ones + zeros > 0) { + /* printf("\t:ones=%d, consensus=%d\n", (int) ones, consensus); */ + /* fflush(stdout); */ + for (j = 0; j < sample_set_size; j++) { + u = sample_set[j]; + if (disagree[u] && (genotypes[j] != consensus) + && (genotypes[j] != TSK_MISSING_DATA)) { + /* This sample has disagreed with consensus twice in a row, + * so remove it */ + /* printf("\t\tremoving %d\n", sample_set[j]); */ + sample_set[j] = -1; + } + } + + /* For the remaining samples, set the disagree flags based + * on whether they agree with the consensus for this site. */ + for (j = 0; j < sample_set_size; j++) { + u = sample_set[j]; + if (u != -1) { + disagree[u] = ((genotypes[j] != consensus) + && (genotypes[j] != TSK_MISSING_DATA)); + } + } + /* Repack the sample set */ + tmp_size = 0; + for (j = 0; j < sample_set_size; j++) { + if (sample_set[j] != -1) { + sample_set[tmp_size] = sample_set[j]; + tmp_size++; + } + } + sample_set_size = tmp_size; + if (sample_set_size <= min_sample_set_size) { + /* printf("BREAK\n"); */ + break; } } } - /* Repack the sample set */ - tmp_size = 0; - for (j = 0; j < sample_set_size; j++) { - if (sample_set[j] != -1) { - sample_set[tmp_size] = sample_set[j]; - tmp_size++; - } - } - sample_set_size = tmp_size; - if (sample_set_size <= min_sample_set_size) { - /* printf("BREAK\n"); */ - break; - } } *last_site_ret = last_site; return ret; diff --git a/lib/tests/tests.c b/lib/tests/tests.c index cda33456..b9e3083f 100644 --- a/lib/tests/tests.c +++ b/lib/tests/tests.c @@ -218,6 +218,52 @@ test_ancestor_builder_multi_site(void) ancestor_builder_free(&ab); } +static void +test_ancestor_builder_stopping_condition(void) +{ + /* + * 8 samples, 5 sites. Site 0 is focal at time 2.0; the rest are older. + * The focal carriers are samples {0, 1}. Sample 1 disagrees with the + * consensus at sites 1 and 3. Site 2 is missing for both carriers and + * must preserve the disagreement. Removing sample 1 at site 3 halves + * the sample set, so extension stops before site 4. + */ + int ret = 0; + ancestor_builder_t ab; + size_t num_samples = 8; + size_t max_sites = 5; + allele_t g0[8] = { 1, 1, 0, 0, 0, 0, 0, 0 }; /* focal */ + allele_t g1[8] = { 1, 0, 1, 1, 0, 0, 0, 0 }; /* first disagreement */ + allele_t g2[8] = { -1, -1, 1, 1, 1, 0, 0, 0 }; /* missing for focal carriers */ + allele_t g3[8] = { 1, 0, 1, 1, 0, 0, 0, 0 }; /* second disagreement */ + allele_t g4[8] = { 1, 1, 1, 1, 0, 0, 0, 0 }; /* beyond the stopping point */ + allele_t ancestor[5]; + tsk_id_t focal_sites[] = { 0 }; + tsk_id_t start, end; + + ret = ancestor_builder_alloc(&ab, num_samples, max_sites, -1, 0); + CU_ASSERT_EQUAL_FATAL(ret, 0); + ret = ancestor_builder_add_site(&ab, 2.0, g0); + CU_ASSERT_EQUAL_FATAL(ret, 0); + ret = ancestor_builder_add_site(&ab, 3.0, g1); + CU_ASSERT_EQUAL_FATAL(ret, 0); + ret = ancestor_builder_add_site(&ab, 3.0, g2); + CU_ASSERT_EQUAL_FATAL(ret, 0); + ret = ancestor_builder_add_site(&ab, 3.0, g3); + CU_ASSERT_EQUAL_FATAL(ret, 0); + ret = ancestor_builder_add_site(&ab, 3.0, g4); + CU_ASSERT_EQUAL_FATAL(ret, 0); + ret = ancestor_builder_finalise(&ab); + CU_ASSERT_EQUAL_FATAL(ret, 0); + + ret = ancestor_builder_make_ancestor(&ab, 1, focal_sites, &start, &end, ancestor); + CU_ASSERT_EQUAL_FATAL(ret, 0); + CU_ASSERT_EQUAL(start, 0); + CU_ASSERT_EQUAL(end, 4); + + ancestor_builder_free(&ab); +} + static void test_ancestor_builder_break_ancestor(void) { @@ -1677,6 +1723,8 @@ main(int argc, char **argv) { "test_ancestor_builder_errors", test_ancestor_builder_errors }, { "test_ancestor_builder_one_site", test_ancestor_builder_one_site }, { "test_ancestor_builder_multi_site", test_ancestor_builder_multi_site }, + { "test_ancestor_builder_stopping_condition", + test_ancestor_builder_stopping_condition }, { "test_ancestor_builder_break_ancestor", test_ancestor_builder_break_ancestor }, { "test_ancestor_builder_one_bit_encoding", test_ancestor_builder_one_bit_encoding }, diff --git a/tests/test_ancestors.py b/tests/test_ancestors.py index 686b38a2..42ef0038 100644 --- a/tests/test_ancestors.py +++ b/tests/test_ancestors.py @@ -236,28 +236,29 @@ def compute_ancestral_states(self, a, focal_site, sites): site_time = self.sites[site_index].time derived_count = self.sites[site_index].derived_count - for j, u in enumerate(sample_set): - if ( - disagree[u] - and (g_l[u] != consensus) - and (g_l[u] != tskit.MISSING_DATA) - ): - sample_set[j] = -1 - if site_time > focal_time: if ones + zeros == 0: a[site_index] = tskit.MISSING_DATA else: a[site_index] = consensus - if (site_time > focal_time) or (derived_count > ones): + if ((site_time > focal_time) or (derived_count > ones)) and ones + zeros > 0: + for j, u in enumerate(sample_set): + if ( + disagree[u] + and (g_l[u] != consensus) + and (g_l[u] != tskit.MISSING_DATA) + ): + sample_set[j] = -1 + for u in sample_set: if u != -1: disagree[u] = ( g_l[u] != consensus and g_l[u] != tskit.MISSING_DATA ) - sample_set = sample_set[sample_set != -1] + sample_set = sample_set[sample_set != -1] + if len(sample_set) <= min_sample_set_size: break