Add Asm2Vec method for comparing traces

This commit is contained in:
Flawed
2023-01-18 19:15:11 -08:00
parent aa0c7500de
commit 7513291e21
5 changed files with 937 additions and 0 deletions
+6
View File
@@ -0,0 +1,6 @@
import importlib
__all__ = ["model", "datatype", "utils"]
for module in __all__:
importlib.import_module(f".{module}", "asm2vec")
+205
View File
@@ -0,0 +1,205 @@
import torch
import random
import warnings
import re
class Token:
def __init__(self, name, index):
self.name = name
self.index = index
self.count = 1
def __str__(self):
return self.name
class Tokens:
def __init__(self, name_to_index=None, tokens=None):
self.name_to_index = name_to_index or {}
self.tokens = tokens or []
def __getitem__(self, key):
if type(key) is str:
if self.name_to_index.get(key) is None:
warnings.warn("Unknown token in training dataset")
return self.tokens[self.name_to_index[""]]
return self.tokens[self.name_to_index[key]]
elif type(key) is int:
return self.tokens[key]
else:
try:
return [self[k] for k in key]
except:
raise ValueError
def load_state_dict(self, sd):
self.name_to_index = sd["name_to_index"]
self.tokens = sd["tokens"]
def state_dict(self):
return {"name_to_index": self.name_to_index, "tokens": self.tokens}
def size(self):
return len(self.tokens)
def add(self, names):
if type(names) is not list:
names = [names]
for name in names:
if name not in self.name_to_index:
token = Token(name, len(self.tokens))
self.name_to_index[name] = token.index
self.tokens.append(token)
else:
self.tokens[self.name_to_index[name]].count += 1
def update(self, tokens_new):
for token in tokens_new:
if token.name not in self.name_to_index:
token.index = len(self.tokens)
self.name_to_index[token.name] = token.index
self.tokens.append(token)
else:
self.tokens[self.name_to_index[token.name]].count += token.count
def precompute_weights(self, pos):
"""
This process actually takes a long time due to the size of the weights,
so precompute them according to the shape of the dataset.
"""
from tqdm import tqdm
total = sum([token.count for token in self.tokens])
token_weights = torch.zeros(len(self.tokens))
for token in self.tokens:
token_weights[token.index] = (token.count / total) ** 0.75
weights = token_weights.repeat(pos.shape[0], 1)
for i in tqdm(range(pos.shape[0])):
for taken in pos[i]:
weights[i][taken] = 0
return weights
def sample(self, all_weights, batch_indices, num=5):
weights = all_weights[batch_indices]
return torch.multinomial(weights, num, replacement=False)
class Function:
def __init__(self, insts, blocks, meta):
self.insts = insts
self.blocks = blocks
self.meta = meta
@classmethod
def load(cls, text):
"""
gcc -S format compatiable
"""
label, labels, insts, blocks, meta = None, {}, [], [], {}
for line in text.strip("\n").split("\n"):
if line[0] in [" ", "\t"]:
line = line.strip()
# meta data
if line[0] == ".":
key, _, value = line[1:].strip().partition(" ")
meta[key] = value
# instruction
else:
inst = Instruction.load(line)
insts.append(inst)
if len(blocks) == 0 or blocks[-1].end():
blocks.append(BasicBlock())
# link prev and next block
if len(blocks) > 1:
blocks[-2].successors.add(blocks[-1])
if label:
labels[label], label = blocks[-1], None
blocks[-1].add(inst)
# label
else:
label = line.partition(":")[0]
# link label
for block in blocks:
inst = block.insts[-1]
if inst.is_jmp() and labels.get(inst.args[0]):
block.successors.add(labels[inst.args[0]])
# replace label with CONST
for inst in insts:
for i, arg in enumerate(inst.args):
if labels.get(arg):
inst.args[i] = "CONST"
return cls(insts, blocks, meta)
def __hash__(self):
return hash("\n".join((str(inst) for inst in self.insts)))
def __eq__(self, other):
if isinstance(other, Function):
a = "\n".join((str(inst) for inst in self.insts))
b = "\n".join((str(inst) for inst in other.insts))
return a == b
return False
def tokens(self):
return [token for inst in self.insts for token in inst.tokens()]
def random_walk(self, num=3):
return [self._random_walk() for _ in range(num)]
def _random_walk(self):
current, visited, seq = self.blocks[0], [], []
while current not in visited:
visited.append(current)
seq += current.insts
# no following block / hit return
if len(current.successors) == 0 or current.insts[-1].op == "ret":
break
current = random.choice(list(current.successors))
return seq
class BasicBlock:
def __init__(self):
self.insts = []
self.successors = set()
def add(self, inst):
self.insts.append(inst)
def end(self):
inst = self.insts[-1]
return inst.is_jmp() or inst.op == "ret"
class Instruction:
def __init__(self, op, args):
self.op = op
self.args = args
def __str__(self):
return f'{self.op} {", ".join([str(arg) for arg in self.args if str(arg)])}'
@classmethod
def load(cls, text):
text = text.strip().strip("bnd").strip() # get rid of BND prefix
text = text.replace(" - ", " + ")
text = re.sub(r"0x[0-9a-f]+", "CONST", text)
text = re.sub(r"\*[0-9]", "*CONST", text)
text = re.sub(r" [0-9]", " CONST", text)
op, _, args = text.strip().partition(" ")
if args:
args = [arg.strip() for arg in args.split(",")]
else:
args = []
args = (args + ["", ""])[:2]
return cls(op, args)
def tokens(self):
return [self.op] + self.args
def is_jmp(self):
return "jmp" in self.op or self.op[0] == "j"
def is_call(self):
return self.op == "call"
+88
View File
@@ -0,0 +1,88 @@
import torch
import torch.nn as nn
bce, sigmoid, softmax = nn.BCELoss(), nn.Sigmoid(), nn.Softmax(dim=1)
# TODO: Document how this shit works because wtf
class ASM2VEC(nn.Module):
def __init__(self, vocab_size, function_size, embedding_size):
super(ASM2VEC, self).__init__()
# Dictionary of the token embeddings: v_t
self.embeddings = nn.Embedding(
vocab_size,
embedding_size,
_weight=(torch.rand(vocab_size, embedding_size) - 0.5) / embedding_size / 2,
)
# Dictionary of the function embeddings, \theta_f_s
self.embeddings_f = nn.Embedding(
function_size,
2 * embedding_size,
_weight=(torch.rand(function_size, 2 * embedding_size) - 0.5)
/ embedding_size
/ 2,
)
# Dictionary of outputs: Transposed v'_t
self.embeddings_r = nn.Embedding(
vocab_size,
2 * embedding_size,
_weight=torch.zeros(vocab_size, 2 * embedding_size),
)
# Where the old embeddings are stored once the training step is done
self.old_embeddings_f = None
def init_estimation_mode(self, function_size_new):
device = self.embeddings.weight.device
embedding_size = self.embeddings.embedding_dim
embedding_weights = self.embeddings.weight
self.embeddings = nn.Embedding.from_pretrained(embedding_weights)
output_weights = self.embeddings_r.weight
self.embeddings_r = nn.Embedding.from_pretrained(output_weights)
self.old_embeddings_f = self.embeddings_f
self.embeddings_f = nn.Embedding(
function_size_new,
2 * embedding_size,
_weight=(
(torch.rand(function_size_new, 2 * embedding_size) - 0.5)
/ embedding_size
/ 2
).to(device),
)
def v(self, inp):
# Retrieve the embeddings for all the context tokens
e = self.embeddings(inp[:, 1:])
# Retrieve the embedding for the function, \theta_f_s
v_f = self.embeddings_f(inp[:, 0])
# Calculate CT(in_(j-1))
v_prev = torch.cat([e[:, 0], (e[:, 1] + e[:, 2]) / 2], dim=1)
# Calculate CT(in_(j+1))
v_next = torch.cat([e[:, 3], (e[:, 4] + e[:, 5]) / 2], dim=1)
# delta(in_j, f_s) = 1/3 * (\theta_f_s + CT(in_(j-1)) + CT(in_(j+1)))
v = ((v_f + v_prev + v_next) / 3).unsqueeze(2)
return v
def forward(self, inp, pos, neg):
device, batch_size = inp.device, inp.shape[0]
v = self.v(inp)
# negative sampling loss
pred = torch.bmm(self.embeddings_r(torch.cat([pos, neg], dim=1)), v).squeeze()
label = torch.cat(
[torch.ones(batch_size, 1), torch.zeros(batch_size, neg.shape[1])], dim=1
).to(device)
return bce(sigmoid(pred), label)
def predict(self, inp, pos):
device, batch_size = inp.device, inp.shape[0]
v = self.v(inp)
probs = torch.bmm(
self.embeddings_r(
torch.arange(self.embeddings_r.num_embeddings)
.repeat(batch_size, 1)
.to(device)
),
v,
).squeeze(dim=2)
return softmax(probs)
+314
View File
@@ -0,0 +1,314 @@
import os
import time
import torch
from torch.utils.data import DataLoader, Dataset
from pathlib import Path
from .datatype import Tokens, Function, Instruction
from .model import ASM2VEC
import re
import json
CONSTANTS_RE = re.compile(r"(-? 0x[0-9a-f]+)|\*([0-9])| ([0-9])")
class TraceData:
"""
A class that stores information about traces read from file.
Terminology:
Trace/Function:
The text produced by the generate_deep_traces.py script for each
opcode switch case.
Token:
A numerical representation of the each symbol that appears in the
trace. The tokens dictionary should be shared across all traces to
ensure the same representation matches.
ID:
A representative identifier for the trace of a switch case. An ID
may point to more than one switch case, if the traces of those
switch cases are textually identical (ignoring constants).
idx:
Index of the ID/Function. This is necessary to index into the
training model embeddings.
Opcode set:
A set of opcodes that a single switch case covers.
Pointer Opcode/ptr_opcode:
A single opcode that represents the opcode set.
Constants vector:
Since the training process removes constants from the input text,
the constants vector is a "signature" generated by looking at
constants used in the trace. These are correlated with ptr_opcodes
and not IDs since constants can differ among cases represented by
the same ID.
"""
def __init__(self, tokens):
self.tokens = tokens
self.traces = dict()
"""
maps trace => ID.
See class docstring for more details.
"""
self.ids = dict()
"""
maps ID => { idx, [ptr_opcodes...] }.
See class docstring for more details.
"""
self.constants_vectors = dict()
"""
maps ptr_opcode => constants_vector.
See class docstriing for more details.
"""
self.opcode_sets = dict()
"""
maps ptr_opcode => [opcodes...]
See class docstriing for more details.
"""
def __process_trace(self, ptr_opcode, text):
fn = Function.load(text)
if fn in self.traces:
id = self.traces[fn]
self.ids[id]["ptr_opcodes"].append(ptr_opcode)
else:
# Use ptr_opcode as an ID
id = ptr_opcode
self.traces[fn] = id
self.tokens.add(fn.tokens())
self.ids[id] = {
"idx": len(self.ids),
"ptr_opcodes": [ptr_opcode],
}
self.constants_vectors[ptr_opcode] = self.__get_constants_vector(text)
def __process_opcode_sets(self, opcode_sets_file):
with open(opcode_sets_file) as f:
data = json.load(f)
self.opcode_sets = {int(op): ops for op, ops in data.items()}
@staticmethod
def __read_trace_from_file(f):
lines = f.readlines()
normalized_lines = []
for line in lines:
normalized_lines.append(" " + line)
# cap lines at 200 since Asm2Vec performance goes way down if
# the text is too long
return "".join(normalized_lines[:200])
@staticmethod
def __get_constants_vector(trace):
constants = []
for line in trace.strip("\n").split("\n"):
match = CONSTANTS_RE.search(line)
if not match or not match.lastindex:
continue
const_str = match.group(match.lastindex)
const_str = "".join(const_str.split())
const = int(const_str, 16)
# We only really care about constants less than 10000
if abs(const) < 10000:
constants.append(const)
return constants
@staticmethod
def load_data(paths, tokens):
"""Reads traces from paths and returns TraceData for that path."""
if type(paths) is not list:
paths = [paths]
filenames = []
for path in paths:
if os.path.isdir(path):
filenames += [
Path(path) / filename
for filename in sorted(os.listdir(path))
if os.path.isfile(Path(path) / filename)
]
else:
filenames += [Path(path)]
trace_data = TraceData(tokens)
for filepath in filenames:
filename = os.path.basename(filepath)
file_split = os.path.splitext(filename)
file_ext = file_split[-1]
if file_ext == ".json" and filename == "opcode_sets.json":
trace_data.__process_opcode_sets(filepath)
elif file_ext == ".asm":
ptr_opcode = int(file_split[0], base=16)
with open(filepath) as f:
text = trace_data.__read_trace_from_file(f)
trace_data.__process_trace(ptr_opcode, text)
return trace_data
class AsmDataset(Dataset):
def __init__(self, x, y):
self.x = x
self.y = y
def __len__(self):
return len(self.x)
def __getitem__(self, index):
return index, self.x[index], self.y[index]
def preprocess(functions, tokens):
x, y = [], []
for i, fn in enumerate(functions):
for j in range(1, len(fn.insts) - 1):
x.append(
[i]
+ [
tokens[token].index
for token in fn.insts[j - 1].tokens() + fn.insts[j + 1].tokens()
]
)
y.append([tokens[token].index for token in fn.insts[j].tokens()])
return torch.tensor(x), torch.tensor(y)
def train(
trace_data: TraceData,
model=None,
embedding_size=100,
batch_size=1024,
epochs=10,
neg_sample_num=25,
calc_acc=False,
device="cpu",
mode="train",
callback=None,
learning_rate=0.02,
):
"""Trains the model on the provided trace data."""
functions = trace_data.traces.keys()
tokens = trace_data.tokens
if mode == "train":
if model is None:
model = ASM2VEC(
tokens.size(),
function_size=len(functions),
embedding_size=embedding_size,
).to(device)
optimizer = torch.optim.Adam(model.parameters(), lr=learning_rate)
elif mode == "test":
if model is None:
raise ValueError("test mode requires a pretrained model")
optimizer = torch.optim.Adam(model.embeddings_f.parameters(), lr=learning_rate)
else:
raise ValueError("Unknown mode")
# Precompute the token weights so that they are cached for later use
inp, pos = preprocess(functions, tokens)
token_weights = tokens.precompute_weights(pos)
loader = DataLoader(AsmDataset(inp, pos), batch_size=batch_size, shuffle=True)
for epoch in range(epochs):
start = time.time()
loss_sum, loss_count, accs = 0.0, 0, []
model.train()
"""
Recall that the model consumes context tokens as input and uses
embeddings for each output token as an output layer of the network. The
model should output a high score for the token that correctly matches
with the given context, and a low score for tokens that don't match the
given context.
Our dataloader outputs samples in the form of
batch_size x (batch_index, input_context, correct_token).
"""
for i, (batch_indices, inp, pos) in enumerate(loader):
for j in range(pos.shape[1]):
# Sample tokens that are not the positive token.
neg = tokens.sample(token_weights, batch_indices, neg_sample_num)
pos_token = torch.unsqueeze(pos[:, j], 1)
loss = model(inp.to(device), pos_token.to(device), neg.to(device))
loss_sum, loss_count = loss_sum + loss, loss_count + 1
optimizer.zero_grad()
loss.backward()
optimizer.step()
if i == 0 and calc_acc:
probs = model.predict(inp.to(device), pos.to(device))
accs.append(accuracy(pos, probs))
if callback:
callback(
{
"model": model,
"tokens": tokens,
"epoch": epoch,
"time": time.time() - start,
"loss": loss_sum / loss_count,
"accuracy": torch.tensor(accs).mean() if calc_acc else None,
}
)
return model
def save_model(path, model, tokens):
torch.save(
{
"model_params": (
model.embeddings.num_embeddings,
model.embeddings_f.num_embeddings,
model.embeddings.embedding_dim,
),
"model": model.state_dict(),
"tokens": tokens.state_dict(),
},
path,
)
def load_model(path, device="cpu"):
checkpoint = torch.load(path, map_location=device)
tokens = Tokens()
tokens.load_state_dict(checkpoint["tokens"])
model = ASM2VEC(*checkpoint["model_params"])
model.load_state_dict(checkpoint["model"])
model = model.to(device)
return model, tokens
def accuracy(y, probs):
return torch.mean(torch.tensor([torch.sum(probs[i][yi]) for i, yi in enumerate(y)]))
def cosine_similarities(model):
"""
Reads the old and new embeddings from the model and returns the pairwise
cosine similarities between these embeddings.
"""
old_f = model.to("cpu").old_embeddings_f
new_f = model.to("cpu").embeddings_f
v_old = old_f(torch.tensor([i for i in range(old_f.num_embeddings)]))
v_new = new_f(torch.tensor([i for i in range(new_f.num_embeddings)]))
cs_matrix = torch.nn.functional.cosine_similarity(
v_old[:, :, None], v_new.t()[None, :, :]
)
return cs_matrix.detach().numpy()