probability maximize expectation

Viewed 291

Given 3n people that the i-th person can pass a test with probability p_i, now you are required to divide them to n groups that each group has 3 people. The score of one group equals 1 if at least two people pass the test, 0 otherwise. In order to maximize the expectation of total score, how do you group them?

I've thought about this problem for a bit, and I think intuitively it makes sense to group two large p_i with a small p_i. Also, i've thought about in the optimal arrangement, swapping any two p_i from different groups should lower the expectation. I can write out mathematically the difference in expectation when swapping two of the students, but it doesn't seem to give any obvious result. I've hit a wall.

1 Answers

Interesting problem. It feels hard to me, since three tends to be the magic number for NP-hardness, and I don't see any kind of convex structure.

I can suggest the following large-neighborhood local search strategy. If we were just trying to match pairs with singles to form groups of three, then the optimal strategy would be to sort the pairs by how likely they are to have exactly one pass, sort the singles by how likely they are to pass, and match them accordingly. To do local search, form initial groups, then repeatedly split each group uniformly at random into a pair and a single, then rematch optimally as above.

Some very rough Python:

import random


def quality(groups):
    return sum(a * b + a * c + b * c - 2 * a * b * c for [a, b, c] in groups)


def main():
    n = 10
    groups = [[random.random() for j in range(3)] for i in range(n)]
    print(groups)
    print(quality(groups))
    for k in range(1000):
        choices = [random.randrange(3) for i in range(n)]
        pairs = [[group[j - 1], group[j - 2]] for (group, j) in zip(groups, choices)]
        pairs.sort(key=lambda pair: pair[0] + pair[1] - 2 * pair[0] * pair[1])
        singles = [group[j] for (group, j) in zip(groups, choices)]
        singles.sort()
        groups = [pair + [single] for (pair, single) in zip(pairs, singles)]
        print(quality(groups))
    print(groups)


main()
Related