Wednesday, April 18, 2012

Digital Extraction and Rotation

This will be a short one. In the Google Code Jam qualification round, Problem C involved rotating the digits of a number to the left, e.g., 12405 -> 24051 -> 40512 -> 5124.

Python lets us do this with easy string <-> int conversions, but these aren't particularly fast. Instead, we can do it mathematically by extracting digits with % and /. This requires us knowing the number of digits in the number beforehand, but this can be calculated with the log10 function (which is relatively slow, because computers store numbers in binary). In the Code Jam problem the number of digits didn't vary, so I could just calculate the length once then pass it in to my shifting function.

>>> def rot(n, length):
     return n % (10**(length-1))*10 + n / (10**(length - 1))

 
>>> rot(12405, 5)
24051
>>> 
>>> rot(24051, 5)
40512
>>> rot(23829239, 8)
38292392

The % / trick is really great, especially in languages such as C where string <-> int conversions are less straight-forward. It works on two principles: integer division with / by a power of 10 allows us to strip off digits from the right. You can visualise this as moving the decimal point to the left, then deleting everything that follows it. For example, 14231 / 102 gives 142.31, which with int division is 142. Similarly, applying the modulo operator % with a power of 10 allows us to strip digits from the right. 14231 % 102 is the remainder when 14231 is divided by 100, which as we saw before is 31.

Can someone please confirm that mathematical shifting is faster than converting it to a string with str() and manipulating that? And also, that log10 is relatively slow?

Hacky but Effective Gradient Descent

Here's an interesting situation I encountered recently:

Given a L of N (1 ≤ N ≤ 100000) integers, where L0..k is sorted in descending order and Lk..N is sorted in ascending order for some secret k, determine the minimum value in the list. Additionally, this list has the property that:
  • In L0..k, if a < b < c then La-Lb ≥ Lb-Lc.
  • In Lk..N, if a < b < c then Lc-Lb ≥ Lb-La.
You may only look at roughly 1000 elements in the list.

The vague bound on permitted number of lookups is due to the way the original problem was phrased. In the original problem, there was a function with discrete domain and range instead of a list, and the function was slow to evaluate. As long as my program ran in under 3 seconds for some test case, it was marked correct for that test case; knowing the speed of the judging computer, this allowed for roughly 1000 function calls.

To rephrase the problem, the list's values form a sort of discrete parabola (though there are no guarantees on symmetry), and we have to find its minimum.

EDIT: I later found out that ternary search was a correct solution. However, one of my friends used the parabola-like property to do a standard binary search over the difference between successive points: i.e., its derivative, (if such a term is defined for this situation), which is a monotonic function. Kudos to Charlie for that beautiful solution).

After trying to code a couple of inevitably buggy ternary searches, I tried a different approach, based on the notion of gradient descent. The idea here is that we have a function whose minimum we want to find, so we pick some starting point on the slope of the function and sort of 'slide' down the slopes (though really, this 'slide' is more of a 'hop' to another point on the function's slope), stopping once we reach the minimum. This is also related to the notion of simulated annealing . When implementing gradient descent with a continuous function, there are two ways of getting the slope of our current point: either differentiate the function at the point (thus giving an exact result) or just approximate it by picking two close points.

I used the latter approach for this problem. We can get the 'tangent' of some point on the list by comparing two adjacent points. For example, if our current position in the list holds the value 1000 and the subsequent position holds the value 500, we should jump forward by a large amount; but if the subsequent position held the value 990, we're probably quite close to the minimum, so we should jump forward by a small amount.

It turns out that 1000 list lookups is a hell of a lot, so even dodgy heuristics like these work most of the time -- it's a number of lookups vs. speed of convergence trade-off, and since our permitted number of lookups is very generous, we can be confident we'll converge to the correct solution.

Pretty much, all I had to do was implement something like this. Note that this specific code is largely untested.

# this list is small, but the solution works on very large lists as well
L = [123912, 12000, 11000, 9000, 5000, 4600, 4300, 4250, 4230, 4228, 4231, 4235, 4250, 5400, 6900, 10000, 234432, 23423432, 8645632423436]

min_point, min_val = -1, float('inf')

# start at list item 0, but really we could start anywhere
cur = 0

# initialise the jump size to some value that's slightly smaller than the size of L
jump = int((len(L)-1 - cur) / 1.5)

for i in xrange(500):
    # this could theoretically go over the length of the list, but that won't happen
    # due to the pattern of jumps. i don't think you can ever normally get to the
    # rightmost element, but this should really be proved
    adj = cur + 1

    # perform two lookups: one for the current point, one for the adjacent point
    val_cur = L[cur]
    val_adj = L[adj]

    # have we found a new best?
    if val_cur < min_val:
        min_point, min_val = cur, val_cur

    if val_adj < val_cur:
        # point to the right is smaller than point to the left...
        # meaning we are on the left of the min, meaning jump right
        cur += jump
    else:
     # otherwise, jump left
        cur -= jump

    # perform a vague reduction of the jump size that will allow it to reach 1 within
    # 500 steps (this should be dependent on list size, but this is all a huge hack).
    # just make sure it doesn't get to 0 prematurely, we always want to allow adjacent jumps!
    jump = max(0, int(jump * 0.8))

# and, magic!
print min_point, min_val

Pretty cool, huh? This is a very crude version of gradient descent, but hey -- it works, and it's surprisingly effective. It found the precise minimum value in all ~30 cases, including ones that were specifically designed to break heuristic programs like this. Sometimes you've just gotta hack something up.

EDIT 2: Here's an implementation of Charlie's algorithm:

# this list is small, but the solution works on very large lists as well
L = [123912, 12000, 11000, 9000, 5000, 4600, 4300, 4250, 4230, 4228, 4227, 4231, 4235, 4250, 5400, 6900, 10000, 234432, 23423432, 8645632423436]
 
min_point, min_val = -1, float('inf')
 
# start at list item 0, but really we could start anywhere
start, end = 0, len(L) - 1
while start <= end:
    mid = (start + end) / 2
    print start, end, mid
    mid_val, next_val = L[mid], L[mid+1]

    if mid_val < min_val:
        min_point, min_val = mid, mid_val

    if next_val > mid_val:
        # we're on the pos grad. slope
        end = mid - 1
    else:
        start = mid + 1

print min_point, min_val

Friday, February 3, 2012

We Need To Talk About Binary Search.

Although the basic idea of binary search is comparatively straightforward, the details can be surprisingly tricky.
-- Donald Knuth

Why is binary search so damn hard to get right? Why is it that 90% of programmers are unable to code up a binary search on the spot, even though it's easily the most intuitive of the standard algorithms?

  • Firstly, binary search has a lot of potential for off-by-one errors. Do you do inclusive bounds or exclusive bounds? What's your break condition: lo=hi+1, lo=hi, or lo=hi-1? Is the midpoint (lo+hi)/2 or (lo+hi)/2 - 1 or (lo+hi)/2 + 1? And what about the comparison, < or ≤? Certain combinations of these work, but it's easy to pick one that doesn't.
  • Secondly, there are actually two variants of binary search: a lower-bound search and an upper-bound search. Bugs are often caused by a careless programmer accidentally applying a lower-bound search when an upper-bound search was required, or vice versa.
  • Finally, binary search is very easy to underestimate and very hard to debug. You'll get it working on one case, but when you increase the array size by 1 it'll stop working; you'll then fix it for this case, but now it won't work in the original case!

I want to generalise and nail down the binary search, with the goal of introducing a shift in the way the you perceive it. By the end of this post you should be able to code any variant of binary search without hesitation and with complete confidence. But first, back to the start: here is the binary search you were probably taught...

Input: sorted list of elements, query term
Output: the index of the first appearance of the query in the list, or an ERROR value otherwise

I propose an alternative definition: a binary search takes as input a (monotonic) function f(x) and a boolean predicate function p(v), and searches over the finite domain of the function for arguments where the predicate is true for the function's value -- i.e., values x such that p(f(x)) is true.

Based on this definition, here are the definitions of the variants:

Upper-bound: find the maximum argument x such that p(f(x)) is true
Lower-bound: find the minimum argument x such that p(f(x)) is true

The traditional binary search described above is a special case of the general lower-bound search, where f(x) = array[x], the domain of f(x) is the set of integers {0, 1, ..., N-1} (N being the length of array) and p(v) = v ≥ query.

In other words, you're searching for the minimum argument x such that array[x] ≥ query. Hey, this is just what we had before!

Let's face it: this description of binary search isn't very helpful. For example, why use this predicate thing when all we need is a simple array[mid] < query in our binary search?

The advantage of this somewhat convoluted definition comes when either the query is not in the array or it's in the array many times. Say you're searching for the first instance of the number 6 in the following array, using the traditional method:

[1, 1, 2, 4, 5, 5, 5, 6, 6, 6, 6, 8, 10, 10, 11]

It should output 7. Let's try this out...

lo, hi = 0, len(arr) - 1
while lo < hi:
    mid = (lo + hi) / 2
    if arr[mid] >= 6:  hi = mid - 1
    else:              lo = mid + 1
print mid  # should print 7

This looks about right. If mid is greater or equal to 6, it'll keep searching to the left, trying to find smaller values of 6. Otherwise, it'll search to the right of mid. Will mid always be the correct index by the end? I run it... apparently not, it outputs 5. Oh, I know! It's because I used ≥ instead of >! Right now it'll keep searching left after it encounters the first 6. OK, change that to a > and run it again. Now I'm getting 9... oh, it must be because I'm accessing arr[mid] at the end instead of arr[lo]. On the final iteration, arr[mid] would have become too large, but arr[lo] will be just right -- it should be at exactly 7, which is what we want. Hit F5; wtf, 10? Undo that ≥/> change from before but keep the other changes to see if it makes a difference -- nope, now 6? And I haven't even considered the issue of inclusive or exclusive bounds...

No joke, I just messed around with the +'s and -'s and bounds and ≥s and >s and ≤s <s and los and mids and his for about 10 minutes and I can't find a single combination which gets me an answer of 7. This is worse than any other bug because you'll undoubtedly end up in a loop of case-bashing: getting it to work for this case, then finding it fails for another, then fixing it for that case and finding it now fails the original case. This style of debugging never works. Instead, turn off your monitor, grab a pen and paper and plan this out.

OK, are we doing upper-bound or lower-bound search? We're finding the minimum index with the value 6; lower-bound then. What's the predicate? It's a lower-bound search, so we want it to return true while we're larger than our desired index and return false when it's smaller. Easy! p(v) = (v >= 6).

So let's do this right. I now write my own p(v) function, even though the logic is ridiculously simple, and I translate the predicate-based binary search definition into code. Most importantly, I introduce a new variable, the best_so_far variable.

def p(v): return v >= 6

lo, hi = 0, len(arr) - 1
best_so_far = None

while lo <= hi:
    x = (lo + hi) / 2
    if p(arr[x]):
        # we found a potential minimum x, but we should still check to see if any smaller ones work
        best_so_far  = x
        hi = x - 1
    else:
        # the predicate is false, so we need to go right to find true values
        lo = x + 1

print best_so_far 

And it works first time. *whistles*

But the problem definition changes! Your boss tells you that now, you must search for the last occurrence of 6!

But hey, that's cool. Re-evaluate the problem: it's now an upper-bound search. Our predicate must return true for values smaller or equal to 6, but start returning false after we get to 7, so the maximum x that p(f(x)) = true is the index of the last 6.

def p(v): return v <= 6

lo, hi = 0, len(arr) - 1
best_so_far = None

while lo <= hi:
    x = (lo + hi) / 2
    if p(arr[x]):
        # we found a potential maximum x, but we should still check to see if any larger ones work
        best_so_far = x
        lo = x + 1
    else:
        # the predicate is false, so we need to go left to find true values
        hi = x - 1

print best_so_far 

Again, it works first go. Notice that I only changed two things: the predicate function (reversed the sign) and the direction we head when we find a predicate=true (i.e. I swapped the lines lo = x + 1 and hi = x - 1). I did not have to mess with any +'s or -'s. No off-by-ones were introduced during the making of this function.

Notice also that I use lo = x + 1 and not lo = x. Similarly, hi = x - 1 and not hi = x. This is a foolproof way to avoid the nasty infinite loop binary search bug caused by integer division -- it ensures that you're never considering a value of x more than once, so you're always narrowing down your search space by at least 1 each time, hence ensuring termination. The use of max/min_so_far gives us complete control over how we're approaching the solution, meaning that we don't need to mess around trying to work out whether it's lo, hi or mid that contain the return value at the conclusion of the algorithm. I personally find that inclusive bounds work best with this form of binary search, but your mileage may vary. If you use exclusive bounds, I make no guarantees on this strategy's correctness.

Yet again, the requirements change. The list is now in descending order, and you need to find the index of the first item less than 5.

[11, 10, 10, 8, 6, 6, 6, 6, 5, 5, 5, 4, 2, 1, 1]

As usual, you only need to make two decisions. Upper- or lower-bound? It's clearly lower-bound: you're finding the first item that satisfies the predicate. What's the predicate? p(v) = (v < 5). Expected output is 11.

def p(v): return v < 5

lo, hi = 0, len(arr) - 1
best_so_far  = None

while lo <= hi:
    x = (lo + hi) / 2
    if p(arr[x]):
        # we found a potential minimum x, but we should still check to see if any smaller ones work
        best_so_far  = x
        hi = x - 1
    else:
        # the predicate is false, so we need to go right to find true values
        lo = x + 1

print best_so_far 

It prints 11.

Do I expect every programmer to write out a trivial p(v) function for every binary search they write? Of course not. It might help you think about the problem, but it's not required. If you take one thing away from this post, let it be this: in any binary search you ever write, whether it be over a list of strings, or a multi-dimensional space, or over the domain of a function that uses the inclusion-exclusion principle on O(1) cumulative sums of rectangular regions of a ternary predicate mapped over the integer values in a 2D grid (this has happened before), you just need to worry about whether it is upper- or lower-bound and what your predicate is.

BAM. No more bugs in a binary search, ever. You can thank me later.

Tuesday, January 3, 2012

Shortest Superstring Problem

If I give you a list of words, can you find the shortest string that contains all the words?

I can!

It turns out that this problem is equivalent to TSP, which means there are no polynomial-time algorithms to solve it. However, we can tackle it the same way as we tackle TSP: do the dynamic programming solution, which although slow, is the optimal correct algorithm. There's a trick to it though, which caught me out the first time: if one word is a complete subset of another and does not appear at the beginning or the end (as in the case ['abc', 'b']), the edge traversal principle doesn't work. In cases like these, the smaller word must first be removed from the list. You can see this below with ['germanic', 'german', 'germ']. EDIT: And if two words are the same, they'll both be substrings of each other, but we don't want to remove both of them. Thanks to Werner Lemberg for picking this up in the comments.

For once, some neat code. What's going on???

all_words = ['ginger', 'german', 'minutes', 'testing', 'tingling', 'minor', 'testicle', 'manage', 'guilt', 'germanic', 'normal', 'malt', 'german', 'germ']
all_N = len(all_words)

words = []

# remove strings which are substrings of others
for i in xrange(all_N):
    for j in xrange(all_N):
        if i != j and all_words[i] != all_words[j] and all_words[i] in all_words[j]:
            break
    else:
        words.append(all_words[i])

N = len(words)
        
 
# determine the numerical overlap between two strings a,b
def overlap(a, b):
    best = 0
    for i in xrange(1, min(len(a), len(b))+1):
        if b.startswith(a[-i:]):
            best = i
    return best
 
 
cost = [[None] * N for _ in xrange(N)]
# Precompute edge costs
# for every pair of words with indices u,v
for u in xrange(N):
    for v in xrange(N):
        # work out the best compressed concatenation you can make with u then v
        cost[u][v] = len(words[u]) + len(words[v]) - overlap(words[u], words[v])
 
                 
cache = {}
backtrace = {}
                 
# top-down DP
def solve(used, last):
    if (used,last) not in cache:
     
        bestCost = 0
        bestOption = None
 
        # the base case is when used == 0. using no words, the optimal solution is of length 0
        if used != 0:
         
            bestCost = float('inf')
         
            # for each word we can use
            for i in xrange(N):
                 
                # if the word is as yet unused
                if (1 << i) & used:
                 
                    # calc the cost of using it
                    newCost = cost[i][last] + solve(used & ~(1<<i), i)
                     
                    # if we've reached a new best solution, update stuff
                    if newCost < bestCost:
                        bestCost = newCost
                        bestOption = i
         
        # cache stuff
        cache[(used, last)] = bestCost
        backtrace[(used, last)] = bestOption
     
    return cache[(used, last)]
 
 
# run it for all possible starting cases
bestCost = float('inf')
bestOption = None
for i in xrange(N):
    cur = solve(((1<<N)-1) & ~(1<<i), i)
    if cur < bestCost:
        bestCost = cur
        bestOption = i
 
 
# reconstruct the words that were used
soln = []
used = (1<<N) - 1
last = bestOption
while last is not None:
    soln.append(words[last])
    used &= ~(1<<last)
    last = backtrace[(used, last)]
soln.reverse()
 
 
# now compress the words of the solution into the final string
cur = soln[0]
for i in xrange(1, N):
    cur += soln[i][overlap(cur, soln[i]):]
print cur

The output is minutestinglingingermanicminormaltmanageguiltesticle. Fascinating stuff, I know.

Saturday, December 24, 2011

Quick 'n Dirty Disjoint Sets

The disjoint-set data structure is magical. Not only is it flexible and indispensable in a variety of situations, but both the associated algorithms and the implementation are remarkably clean and simple. This makes the disjoint-set data structure (sometimes called the 'union-find data structure', named after its two primary operations) a joy to encounter when programming.

The data structure is used to answer the question "given a set of bidirectional connections between nodes, can I reach node b from node a by walking along these connections?" It can be visualised as a set of disconnected trees where each tree regularly bumps deep nodes up to the top, where they become children of the root node. Each tree represents a set of connected nodes, so to determine whether we can walk between two nodes -- that is, whether the two nodes are part of the same set -- all we have to do is check if they have the same root. It's marginally more complicated than that, but that's the general gist of union-find. The data structure also allows quick merging of trees, meaning you can add new connections in near-constant time.

Being a tree data structure, it's relatively easy to implement using a pointer structure, where each node points to its parent. Here's a C implementation:

struct set {
    int rank;
    struct set *parent;
};

struct set* newSet(void) {
    struct set *s = (struct set*) malloc(sizeof(struct set));
    s->parent = s;
    s->rank = 0;
    return s;
}

struct set* find(struct set *s) {
    if (s->parent == s)  return s;
    else                 return (s->parent = find(s->parent));
}

void join(struct set *a, struct set *b) {
    struct set *aRep = find(a), *bRep = find(b);
    if (aRep->rank > bRep->rank) {
        bRep->parent = aRep;
    } else if (aRep->rank < bRep->rank) {
        aRep->parent = bRep;
    } else {
        aRep->parent = bRep;
        bRep->rank ++;
    }
}

It's short, sweet and understandable (I'm glaring at you, binary index trees): exactly the kind of thing you can code up in a programming competition. The rank property is used to implement union by rank, optimising the data structure's performance. Combined with path compression, this yields a time complexity of O(inverse Ackermann(n)) for all operations, which is below 5 for all practical values of n and so is effectively a constant time complexity.

But hey, that's still pretty long, and it needs, like, pointers and stuff. Why don't we drop union by rank? Although we lose the inverse Ackermann upper bound, path compression by itself keeps operations almost as fast... and who could resist the reduction in code size?

rep = range(0,100)

def find(p):
    if rep[p] != p:
        rep[p] = find(rep[p])
    return rep[p]
 
def union(p,q):
    rep[find(p)] = find(q)

This is almost the same thing, but with less code and more elegance (in my opinion, at least). I've been told that at least one university teaches union-find without union by rank, so it's also an acceptable implementation.

Here, I'm generating a set of random points on the plane and building a Euclidean minimal spanning tree (MST) from them, using Kruskal's algorithm and our very own disjoint-set data structure implementation. This code has not been tested for correctness.

import random, itertools

rep = range(0,100)

def find(p):
    if rep[p] != p:
        rep[p] = find(rep[p])
    return rep[p]
 
def union(p,q):
    rep[find(p)] = find(q)


# generate 100 random points on the plane
pts = [(random.randint(0,1000), random.randint(0,1000)) for i in xrange(100)]

# define Euclidean distance function
def dist(a,b):
    return ((pts[a][0]-pts[b][0])**2 + (pts[a][1]-pts[b][1])**2)**0.5

# generate pairs of points representing the edges of the complete graph
edges = list(itertools.combinations(range(100), 2))

# sort the edges by length, in ascending order
edges.sort(cmp=lambda (u1,v1),(u2,v2): -1 if dist(u1,v1) < dist(u2,v2) else 1)

# Kruskal's algorithm for MST
total_cost = 0
for (u,v) in edges:
    if find(u) != find(v):
        total_cost += dist(u,v)
        union(u,v)

print total_cost

It might not be as theoretically sound as the pointer implementation, especially considering that the max number of nodes is also constrained by the size of the rep array, but it's quick and dirty, and sometimes that's just what you need.

Tuesday, October 11, 2011

ASCII Circles

While browsing Reddit, I came across this page. It's a fun little task. Here's my solution, with lots of fiddly maths to make the circle look as clean as possible.

import sys
r = int(sys.argv[1])

for y in xrange(-r*1.3,r*1.3,2):
    for x in xrange(-r*1.3,r*1.3):
        sys.stdout.write('#$@%&0*;:,. '[min(abs(r*r - (x*x+y*y))/(r/2), 11)])
    print

Circle of r = 10:

        ,;*0&&&0*;,       
     ,*%$#@%%%%%@#$%*,    
   ,*@#%0;:,...,:;0%#@*,  
  ,0$@0;.         .;0@$0, 
 .*$@0:             :0@$*.
 ,&#%;.             .;%#&,
 ,&#%;.             .;%#&,
 .*$@0:             :0@$*.
  ,0$@0;.         .;0@$0, 
   ,*@#%0;:,...,:;0%#@*,  
     ,*%$#@%%%%%@#$%*,    
        ,;*0&&&0*;,      

r = 30:

                             .,:;*00&&&&&&&00*;:,.                            
                         ,;0&@$##$@@%%%%%%%@@$##$@&0;,                        
                     .;0%$#$%&*;:,,..     ..,,:;*&%$#$%0;.                    
                   :0%#$%0;:.                     .:;0%$#%0:                  
                .;&$#%0;,                             ,;0%#$&;.               
               ;&$$%*,                                   ,*%$$&;              
             ,0@#%*,                                       ,*%#@0,            
            ;&#@0:                                           :0@#&;           
           ;%#%*,                                             ,*%#%;          
          ;%#%;.                                               .;%#%;         
         :&#%*.                                                 .*%#&:        
        ,0$@*,                                                   ,*@$0,       
        ;%#&:                                                     :&#%;       
       ,0$@*,                                                     ,*@$0,      
       :&#%;.                                                     .;%#&:      
       :&#%;                                                       ;%#&:      
       :&#%;                                                       ;%#&:      
       :&#%;.                                                     .;%#&:      
       ,0$@*,                                                     ,*@$0,      
        ;%#&:                                                     :&#%;       
        ,0$@*,                                                   ,*@$0,       
         :&#%*.                                                 .*%#&:        
          ;%#%;.                                               .;%#%;         
           ;%#%*,                                             ,*%#%;          
            ;&#@0:                                           :0@#&;           
             ,0@#%*,                                       ,*%#@0,            
               ;&$$%*,                                   ,*%$$&;              
                .;&$#%0;,                             ,;0%#$&;.               
                   :0%#$%0;:.                     .:;0%$#%0:                  
                     .;0%$#$%&*;:,,..     ..,,:;*&%$#$%0;.                    
                         ,;0&@$##$@@%%%%%%%@@$##$@&0;,                        
                             .,:;*00&&&&&&&00*;:,.                            
                                                                  

The explanation I posted on Reddit (sorry about the formatting, hopefully it's still clear).

I guess I should explain what this does. Like the challenge poster, I use the Pythagorean theorem. I look through each coordinate in a square of size 2.6r * 2.6r around the centre-point. The circle only really exists in the 2r * 2r square around the point, but the extra padding on all sides allows for the anti-aliasing to continue further along the edges.

Mathematically, a circle consists of the coordinates (x, y) that satisfy x2 + y2 = r2, r being the radius of the circle. Any points we encounter that satisfy this should be "coloured in" as dark as possible. Some points do not fall exactly on the circle, but they fall close to the circle: the (absolute) difference between r2 and x2 + y2 is (if I am not mistaken) the distance between (x,y) and the closest point on the circle to (x,y). If the difference is 0, we have the case above. Otherwise, the larger the value, the less it is part of the circle and so the lighter we should 'colour' it.

To colour points, I index the 'colouring' string by a scaled value of this closeness, scaled so the width of the ring remains the same for different values of r.

Thursday, October 6, 2011

Space-efficient Memoization with dicts

Python's not usually considered a good language for memory- or time-intensive tasks, but its extensive standard library and variety of built-ins make it an invaluable tool for algorithmic work.

Consider the task of memoizing a recursive function. Typically, a large d-dimensional array is used as a cache, where d is the number of parameters in the subproblem's state. This is great for bottom-up iterative dynamic programming, where the optimal solution for every state must be evaluated -- every cell of the array will be occupied, so there's no redundant space. However, many problems do not require that every state be examined. Consider a case of the unbounded knapsack problem where you have a knapsack of size 15 and 3 items available of costs 5, 7, and 11. The bottom-up approach will evaluate the optimal solution for states 1, 2, 3, 4, 6, 7, 8 and 9, none of which contribute at all to the final optimal selection of items; after all, there's no way to fill a knapsack of those sizes given those items. The top-down memoized recursion approach, on the other hand, will only evaluate states which can theoretically contribute to the optimal solution (even if they do not).

Despite the overhead created by recursive calls and a much higher proportion of cache misses, the top-down memoized approach is often a lot faster than the bottom-up dynamic programming solution. However, people will still generally use the d dimensional array for caching, even when memoizing. This leads to lots of unused cells and wasted space.

Python gives you another option: instead of an array, use a dictionary of {d-tuple representing state : optimal solution for state}. Dictionaries have extremely fast look-up times, and the amount of memory saved is substantial.

Consider the 'matrix sum' problem (as seen in Project Euler #345), a variant of the classic n-rooks problem. You're given a n x n grid (chessboard?) where each cell has a value associated with it. Your task is to place n rooks on the board such that none of them threaten another (no two rooks lie on the same row or column) and the sum of values in the cells they occupy is maximised.

This is an exponential dynamic programming problem, where the state of a subproblem has two parameters: the current row, and a binary string representing the columns that have already been occupied and cannot have rooks placed in them from now on. I don't want to elaborate too much on the solution because it's an interesting problem that's worth your attention. There are n possible values for the current row and 2n possible values for the binary string, so in total there should be n * 2n states. Let's work with n = 15. The total theoretical number of states should be 491,520, which would be the number of cells in the d-dimensional array.

However, in practice there are many states that are impossible to reach. For example, there is no state where you're on the very first row and yet all columns are already occupied. Similarly, there's no state where you're on the last row and no columns are occupied. It's possible to mathematically calculate the exact number of reachable solutions (I did it once, it's not very fun) but since this is a programming blog, I'd rather just demonstrate a solution that tells you exactly how many states must be cached.

N = 15
m = [map(int, line.split()) for line in open('matrix.txt')]

cache = {}
def dp(r, u):
    if (r,u) not in cache:
        cache[(r,u)] = 0 if r==N else max(dp(r+1, u|(1<<c)) + m[r][c] for c in xrange(N) if not u&(1<<c))
    return cache[(r, u)]

print dp(0,0), len(cache)

Output:

32768

That's right ladies and gentlemen, only 32768 states are reachable. That's 1/15th of the original memory usage. Obviously you could do the same thing in another language, but Python makes it so easy. Have fun implementing your own hash table in C :)

Monday, September 5, 2011

A Practical Use for the 'complex' Type

In a recent programming competition, a question asked competitors to generate certain fractals (Julia sets, to be precise). This task requires understanding of the mathematics behind fractals; specifically, complex numbers. The judges assumed that most competitors would be unfamiliar with complex numbers and how to manipulate them programmatically. Hence, about half of the problem statement documented the different operations that could be done with complex numbers: polar/Cartesian conversions, getting the modulus and the angle, how the arithmetic operators worked, etc.

Luckily, I recalled that Python has built-in support for complex numbers. All I had to do was plug in the given formulae and the problem basically solved itself.

In celebration, here's a short piece of Python code which generates a 200x200 Mandelbrot set in ASCII, using Python's complex number support and the cmath library. cmath is to complex numbers as math is to real numbers: it provides loads of helpful functions you might need when dealing with complex numbers. Here, I use cmath to get the modulus of z.

import math, cmath
RX, RY = (-2.5,1), (-1,1)

# maps a pixel (x,y) into a (x,y) point on the Mandelbrot canvas
def lmap(n, (a,b)): return n/200.0 * (b-a) + a

def value(x,y):
    z, c = 0, complex(lmap(x, RX), lmap(y, RY))
    iters = 0
    while cmath.polar(z)[0] < 4 and iters < 1000:
        z = z*z + c
        iters += 1
    return int(math.log(iters)/math.log(1000) * 10.0)
        
out = open('out.txt', 'w')
for r in xrange(200):
    for c in xrange(200):
        out.write(' .:;*0&%@$#'[value(r,c)])
    out.write('\n')

A small portion of the generated fractal:

Wednesday, August 3, 2011

Anagram Searching -- A Job Interview Question

Occasionally a problem comes along that stuns me in its simplicity and elegance. This one in particular is from the 2006 International Olympiad in Informatics (IOI) under the name Writing. It's a great question to give in a job interview because a bruteforce solution is easy to work out, and some may stop there. However, with a bit more effort, it's possible to develop a solution that improves on the bruteforce by orders of magnitude.

A summary of the problem is as follows:

You're given a parent string and a query string of assorted characters of some alphabet (ASCII works nicely). You must determine how many times the query string -- or a permutation of the query string -- appears in the parent string.

Here is a typical thought-path one may take to the optimal solution. In order to evaluate the efficiency of each algorithm, we'll call the parent string's length n, the query string's length k and the alphabet size s.

1. Bruteforce


Generate all permutations of the query string, then do a naive sub-string count on the parent string for each permutation.

import itertools
sum(parent.count(perm) for perm in itertools.permutations(query))

Complexity: O(k!) for generating the permutations, then O(nk) for a naive substring counting algorithm, giving O(nk * k!).


2. "Crossing off" method


There are lots of common string search/manipulation problems in circulation, and most involve strings that must be treated as ordered sequences; that is, you can't mess around with the order of the characters.

Instead of treating this problem as an extension of those problems with an extra restriction, take advantage of the unordered property.
Here's a new way of looking at the problem: for every consecutive set of k characters in the parent string, is that set of k characters the same as the set of characters in our query string?

We've now introduced a new problem: efficient comparison of character sets.

Let our first string be "goldy" and our second be "dogle". A fairly obvious way of checking that the characters are the same is by using a "crossing-off" method. For each character in "goldy", check if it's in "dogle". If it is, and you haven't encountered it on a previous pass, cross it off. make sure the character you find in the second string not already been crossed off; "aab" should not be considered the same as "abb". If at the end of the process all characters have been crossed off and early termination didn't occur then your strings match.

Complexity: For every consecutive set of k characters in the parent (O(n) sets), for every character in that set, check if it's in the query string and if so, cross if off. This algorithm has a time complexity of O(nk2) which is already a massive improvement over the brute-force above.

3. Sort + compare


We can improve on the cross-off method. A better way to compare two sets is to sort both sets then perform an element-by-element comparison. "goldy" becomes "dgloy", "dogle" becomes "deglo", and a standard string comparison shows that the strings are unequal. Since you're probably in a language with a linearithmic sort in its standard library, this solution is easy to implement and works well.

sortedQuery = sorted(query)
print sum(sorted(parent[i:i+len(query)]) == sortedQuery for i in xrange(len(parent) - len(query) + 1))    

Complexity: You're still going through every consecutive block of k characters in the parent string, but this time you just have to compare the sorted block to the sorted query string. O(nk lg k).

4. Constant-time lookup.


It turns out that sort-based set comparison is suboptimal. We can take advantage of the array structure's constant-time random access to create a lookup table of the form lookup[char] = number of occurences of 'char' in query string, for every value of char in our alphabet. We first generate the lookup table in O(k) time and store it. As usual we'll iterate through every consecutive block of k characters in the parent string, but this time instead of sorting, we populate a second lookup table with the data from this block. To compare the sets all we have to do is compare the lookup tables.

Complexity: For each block, generate lookup tables (O(k)) then compare (O(s)). O(n(s + k)).

There's a slight optimisation to be made. If s is large and k is small, our table comparison will end up going through lots of characters which never appear in the query. To prevent this, we can store the characters we come across when constructing the second lookup table. When we compare the tables, we only have to compare the counts for these characters.

import string

numMatches = 0

# allows for the O(nk) optimisation
queryChars = set(query)

correctTable = [0] * 127
for c in query:
    correctTable[ord(c)] += 1

testTable = [0] * 127

for i in xrange(len(parent)-len(query) + 1):
    thisBlock = parent[i:i+len(query)]
    
    for c in thisBlock:
        testTable[ord(c)] += 1

    for c in queryChars:
        if correctTable[ord(c)] != testTable[ord(c)]:
            break
    else:
        numMatches += 1
        
    for c in thisBlock:
        testTable[ord(c)] -= 1

print numMatches

Complexity: Generate the second lookup table in O(k), compare them in O(k) too. O(nk).

5. Magic


The dominating mind-frame when one thinks about this problem is that it's just a twist on the well-known problem of set comparisons: the solution has to loop over each consecutive k-block in the parent string, and my job is to make the comparison between the block and the query string as fast as possible. The block string/query string comparison is lower-bounded by O(n(s + k)).

But this approach always leaves with you two loops, even though the inner loop covers the exact same content as the outer loop. The optimal solution is obtained by eliminating redundancy in the comparison loop by integrating it with the outer loop.

We have to have an O(n) somewhere, if only for inputting the parent string. Our previous solution was doing two O(k) things on top of that: generating the second lookup table, then comparing the two tables. Say we generate a lookup table for the first block of the parent string, finish doing all our processing for that block, then move on to our next block. Our previous algorithm will generate a brand new lookup table now, even though we've just shifted over by a single character. The only rows that can look different between the two tables are the row of the first character of the previous block (which is now not covered by our new block) and the row of the last character in our new block (which was not previously covered). All the other rows will look exactly the same each time. Now that we're only updating two entries each time, generating the table is now O(k) for the first iteration and O(1) for each of the n-k blocks following: O(n). We just stumbled upon an O(n * min(k,s)) solution (you can derive the complexity from the observations above), but let's go straight on.

Now to apply this logic to comparing the two lookup tables. Previously, it was easy to know whether or not a given block's table matches our query string's table: if all of the character counts were the same, they match. The goal is to only operate on the character our block 'left behind' and the character our block is 'picking up'.

Throughout the program, keep track of a distance (similar to the concept of Hamming distance). distance will contain an integer representing the number of unique characters in our current block's table whose counts match that character's count in the query string's table. This will become clearer in a moment. Consider the query string 'hello and the parent string 'ohelol':

current block: 'ohelo'
h e l o
query string table 1 1 2 1
current block table 1 1 1 2
matches? 1 1 0 02

Here, the distance is 2 because only two character counts -- that of the 'h' and that of the 'e' -- match. Now look at the next block:

current block: 'helol'
h e l o
query string table 1 1 2 1
current block table 1 1 2 1
matches? 1 1 1 14

Now our distance is 4. If this distance is equal to the number of unique characters in the query string (which it is), our current block is a matching block. Now do what you were doing before, but instead of comparing the lookup tables, just modify the distance based on whether the 'new' block character and 'old' block character are contained inside the query string, then compare the distance to the number of unique characters in the query string. Phew!

Complexity: O(n). Can I go to sleep now?

This is an incredibly irritating algorithm to code. Here's a C++ implementation (fitting to the 'horrendous style' paradigm of this blog) to keep you busy.

for (int i=0 ; i<query_length ; i++) {
	scanf("%c", &temp);
	query_table[letter_index(temp)] ++;
}
scanf("\n%s", parent);

for (int i=0 ; i<=letter_index('z') ; i++) expected_distance += (query_table[i] > 0);

for (int i=0 ; i<parent_length ; i++) {
	if (i - query_length >= 0) {
		distance -= block_table[letter_index(parent[i-query_length])] == query_table[letter_index(parent[i-query_length])];
		distance += --block_table[letter_index(parent[i-query_length])] == query_table[letter_index(parent[i-query_length])];
	}
	distance -= block_table[letter_index(parent[i])] == query_table[letter_index(parent[i])];
	distance += ++block_table[letter_index(parent[i])] == query_table[letter_index(parent[i])];
	output += (distance == expected_distance);
}



Friday, June 17, 2011

A Peculiar Numbering System

Using the partial function application and composition from the previous post, we can represent the natural numbers in an interesting way.

def partial(f, *p):
    return lambda *q: f(*(p + q))

def compose(*fs):
    return lambda x: reduce(lambda a,b: b(a), (fs+(x,))[::-1])

def represent(n):
    return compose(*(partial(int.__add__, 1),) * n)(0)

assert represent(0) == 0
assert represent(1) == 1
assert represent(1337) == 1337

The nth natural number is the function λx.x+1 composed with itself n times, applied onto the number 0.

This is related to the concept of Church numerals.

Friday, March 4, 2011

Partial Function Application and Function Composition

This following function demonstrates one of my favourite concepts in functional programming -- that of partial function application:

def pfa(f, *p):
    return lambda *q: f(*(p + q))

Partial function application is where you fix particular parameters of a function, producing a new function with the fixed parameters substituted in. For example, say I had a function add(a,b) = a + b. I can use partial function to 'fix' the first argument, a, to a particular number x, and this produces a new function add(b) = x + b.

In practice, PFA is often a more compact and readable version of the ugly anonymous functions you give to higher-order functions such as lambda. For example, consider the following code extract which transforms [a,b,c,...] to [2^a, 2^b, 2^c, ...]:

map(lambda n: pow(2, n), list)

Now take a look at how it is done with PFA:

map(PFA(pow, 2), list)

PFA in this instance creates a new function by substituting the value '2' as the first argument of the pow function, transforming pow(b, e) = b^e into pow(e) = 2^e.

Haskell's PFA notation is even more elegant; because everything in Haskell is just chained evaluation of partial functions, there's no special or unique notation for it, so the following is valid code:

map (2^) list

That maps the partial function 2^ onto list. Haskell also has lambdas, but they're no where near as short or understandable.

map (\x -> 2^x) list

The Python implementation above can also fix multiple arguments at one time, as in pfa(f(a,b,c), a, b) -> f(c). Sometimes it may be preferable to fix the second or third argument, while leaving the first argument variable -- Haskell deals with this elegantly as usual, but I haven't found a nice way to do it in Python.

note: apparently the functools module already has a partial() function for PFA. whoops :)

Function composition is another concept (again, originally mathematical) with similar roots. Here's an implementation:

def compose(*fs):
    return lambda x: reduce(lambda a,b: b(a), (fs+x)[::-1])


Essentially, this takes two functions f(x) and g(y) and produces a new function h(y) == f(g(y)). The new function is often notated as (f.g)(y), as in "f.g is the composition of f and g." Like partial function application, it is rarely a necessity to use function composition, but again it can make things more compact and more intuitive.

A common construct I use when dealing with strings is as follows:

new_s = old_s.strip().split()

This can be composed as follows:

new_s = compose(str.split, str.strip)(old_s)

A better example:

hash(bin(bool(id(sum(range(10))))))
# becomes
compose(hash, bin, bool, id, sum, range)(10)

Sunday, October 10, 2010

Balanced Brackets

Let us define a simple recursive language:

balanced = '()' | '(', balanced, ')' | balanced, balanced;

The theorems (valid clauses) of our language are made up only of the parentheses. We will call a string 'balanced' if and only if it can be expressed in our language. A friendlier definition of a balanced string is as follows:

  • () is a balanced string
  • Putting a balanced string inside a pair of opening and closing parentheses forms a balanced string
  • Concatenating two balanced strings forms a balanced string

Here are some examples:

Balanced
()
(())
()()
(()())
((()(()))()(()()))
((())(()))((())(()))
Not balanced
(
)
((
))
)(
(()
())
(()()(())

The task is to write a program which determines whether a given string is balanced. Additionally, the algorithm must reflect the language's definition as closely as possible. To reflect this goal I'm going to propose an arbitrary restriction: the algorithm should not involve any change in state. The only exception is the string being tested, which may change. Essentially what this means is that you can't use any variables which change their value. The algorithm will also almost certainly be recursive -- which makes sense because the language is defined recursively as well.



This took me a frustratingly long time. I rewrote the code several times over, each time slightly changing my algorithm to make it cleaner and more logical.

My function takes two parameters: a 'string' to test' (which is actually a list of that string, string) and a character which, when found at the same recursive depth, will signify that the string is balanced (expected).

The first thing in the function is a loop for handling cases where there are two or more balanced strings sitting at the same recursive depth.

I get the first element of the string (while at the same time removing it -- str.pop has that side effect) and store it in the variable first. I then do three tests:

  • If first is an opening parenthesis, I check if the string following it is balanced by recursively calling the function on the rest of string. I also tell it to stop when it reaches a ) character on the same recursive depth. If it tells me that it's not balanced, I stop and return False (if the substring isn't balanced then the whole string isn't balanced).
  • If first == expecting, this string (possibly a substring) is balanced because we've reached the closing character on the same recursive level as the opening one and so everything in between must have been balanced.
  • If anything else has happened, the string isn't balanced and we return False.

The loop can only quit normally (without getting out via a return statement) in two ways:

  • The string ended prematurely. In this case, we will still be expecting something, so the expression 'not expecting' will return False. Since the string ended prematurely the string isn't balanced, so we return False correctly.
  • The string ended at the top recursive level as it should. We aren't expecting anything, so it will return True.


def balanced(string, expecting=False):
    while string:
        first = string.pop(0)
        if first == '(':
            if not balanced(string, ')'):
                return False
        elif first == expecting:
            return True
        else:
            return False
    return not expecting

And my testing:

assert balanced(list('((()))'))
assert balanced(list('()'))
assert balanced(list('()()'))
assert balanced(list('(())((()))'))
assert balanced(list('()(()()(())()((())()))'))
assert balanced(list('(((((((()))())())())()))'))
assert balanced(list('((()(())())(())()(()()))'))
assert not balanced(list(')('))
assert not balanced(list('(('))
assert not balanced(list('()))'))
assert not balanced(list('(((()))'))
assert not balanced(list(')))'))
assert not balanced(list('(()'))
assert not balanced(list('(()()()'))