File size: 2,548 Bytes
e8b8483
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
"""Emit one Verilog module per rule in rules.json.

    python rtl_gen.py

Inputs are the selected channels as signed INT8, post-LayerNorm and post-max-pool.
Output is one bit. Combinational, no multipliers, no memory, no constants: the
comparison is between the two sums, so there is nothing to bake in.
"""
import argparse
import math
from pathlib import Path

from common import read_artifact

HERE = Path(__file__).resolve().parent
HEADER = '''// Zero-parameter person classifier, {n} dims.
// Generated by rtl_gen.py; do not edit by hand.
//
// sum(pos) > sum(neg), on signed INT8 channels taken from the pooled feature
// vector. No threshold, so no constant appears anywhere in this module.
'''


def _decl(dims, per_line=8):
    rows = []
    for i in range(0, len(dims), per_line):
        names = ', '.join(f'f{d}' for d in dims[i:i + per_line])
        rows.append(f'    input  signed [7:0] {names},')
    return '\n'.join(rows)


def _sum(dims, width, per_line=4):
    pad = width - 8
    terms = [(f'{{{{{pad}{{f{d}[7]}}}}, f{d}}}' if pad else f'f{d}') for d in dims]
    rows = [' + '.join(terms[i:i + per_line]) for i in range(0, len(terms), per_line)]
    return ' +\n        '.join(rows)


def emit(name, pos, neg):
    k = max(len(pos), len(neg))
    width = 8 + max(1, math.ceil(math.log2(k))) if k > 1 else 8
    body = f'''
module {name} (
{_decl(list(pos) + list(neg))}
    output person_present
);
'''
    if k == 1:
        body += f'    assign person_present = f{pos[0]} > f{neg[0]};\nendmodule\n'
        return HEADER.format(n=len(pos) + len(neg)) + body
    body += f'''    wire signed [{width - 1}:0] pos_sum =
        {_sum(pos, width)};

    wire signed [{width - 1}:0] neg_sum =
        {_sum(neg, width)};

    assign person_present = pos_sum > neg_sum;
endmodule
'''
    return HEADER.format(n=len(pos) + len(neg)) + body


def generate(out_dir=None, rules_json=None):
    out_dir = Path(out_dir or HERE / 'rtl')
    out_dir.mkdir(parents=True, exist_ok=True)
    doc = read_artifact(rules_json or HERE / 'rules.json')
    written = []
    for name, r in doc['rules'].items():
        path = out_dir / f'person_{name}.v'
        path.write_text(emit(f'person_{name}', r['pos_dims'], r['neg_dims']),
                        encoding='utf-8')
        written.append(path)
    return written


if __name__ == '__main__':
    ap = argparse.ArgumentParser(description=__doc__)
    ap.add_argument('--out', type=Path, default=None)
    args = ap.parse_args()
    for p in generate(args.out):
        print(f'wrote {p}')