関数呼び出し
関数呼び出し
自然言語のトレーディング要求を、関数名と閉じた集合の引数を信頼度付きの TypeSafe の質問に対応付けて、通常の型付き関数の呼び出しに変換します。
「ラージのアイスオーツラテ、シュガーなし」と注文すると、バリスタはその文を書き留め たりしません。カップに 4 つの選択肢を印で示します。この cookbook はトレーディング API に ついて同じことをします。文を入力すると、関数名とその引数が評価済みの enum として出て きます。それぞれに信頼度が付きます。
"plot rolling correlation between nvda and spy for the past month"
rolling_correlation(symbol='NVDA', benchmark='SPY', window='1mo') confidence 0.91
"compare nvda amd and msft over the past three months"
compare_returns(symbols=['NVDA', 'AMD', 'MSFT'], window='3mo') confidence 0.94
"show me apple daily with volume"
plot_price(symbol='AAPL', resolution='1d', include_volume=True) confidence 0.75
"what tickers do you have"
list_symbols() confidence 1.00
これらの呼び出しは、トレーディングアシスタント内の 10 個の通常の関数に届きます。引数は
固定リストから値を取るので、すでに Literal です。
def plot_price(
symbol: Literal["SPY", "NVDA", "AMD", "AAPL", "MSFT", "TSLA"],
style: Literal["line", "candles"] = "line",
resolution: Literal["1m", "5m", "15m", "1h", "1d"] = "15m",
window: Literal["1d", "1w", "1mo", "3mo"] = "1w",
include_volume: bool = False,
moving_average: Literal["9", "20", "50"] | None = None,
log_scale: bool = False,
): ...
値が固定リストから来る引数は閉じた集合です。そのリストから 1 つの値を取るとき、まさに
その値たちに対する Choice の質問が与えられるので、関数に届くものは関数が受け付ける値
です。関数自体には手を付けません。追加するのは、各引数が何を意味するかを平易な言葉で
記した spec です。最終的には、自分の関数に向けられる Dispatcher が得られます。
準備
pip install ipython polars matplotlib numpy 'cooksafe>=0.2.0,<0.3.0'
TYPESAFE_API_KEY を設定します。このファイルの隣には 2 つのモジュールがあります。
trader.py は 10 個の関数と、キャッシュから答えを読む TypeSafe クライアントを保持し、
再描画すると API を呼ばずに以下の数値を再生します。dispatch.py は、シグネチャと spec を
読み取って呼び出しを行うコードを保持します。
import json
from pathlib import Path
from cooksafe import make_playground_link
from dispatch import ROUTE, Dispatcher, closed_sets
from IPython.display import Markdown, display
from trader import TOOLS, client, load
TYPESAFE_MODEL = "jev-1.12"
print(f"{len(TOOLS)} functions over {load().height:,} one-minute bars")
10 functions over 156,780 one-minute bars
シグネチャ内の閉じた集合を見つける
型ヒントは、どの引数が固定リストから来るか、各リストに何が入っているかをすでに示して
います。closed_sets はシグネチャを読み取り、それらの引数を 3 つの形に分類します。
choice(Literal なので、リストから 1 つの値)、set(list[Literal[...]] なので
任意の個数)、flag(bool なのでオンかオフ)です。10 個の関数はすべて trader.py に
定義されています。
for name, fn in TOOLS.items():
shapes = closed_sets(fn)
print(
f" {name:<20}{len(shapes)} "
+ ", ".join(f"{a}:{s}" for a, (s, _) in shapes.items())
)
print(
f"\n{sum(len(closed_sets(fn)) for fn in TOOLS.values())} fillable arguments in total"
)
list_symbols 0
market_summary 1 window:choice
plot_price 7 symbol:choice, style:choice, resolution:choice, window:choice, include_volume:flag, moving_average:choice, log_scale:flag
intraday_pattern 3 symbol:choice, window:choice, metric:choice
compare_returns 3 symbols:set, window:choice, normalize:flag
rolling_correlation 4 symbol:choice, benchmark:choice, window:choice, resolution:choice
summary_stats 2 symbol:choice, window:choice
volatility 3 symbol:choice, window:choice, annualized:flag
top_movers 2 window:choice, direction:choice
drawdown 3 symbol:choice, window:choice, plot:flag
28 fillable arguments in total
top_movers は何が除外されるかを示します。3 つの引数のうち 2 つが閉じた集合です。3 つ目の
limit は int なので、質問は与えられず、デフォルトの 3 を保ちます。自由テキスト、
数値、日付も同じです。質問はなく、関数のデフォルトがそのまま使われます。
spec を書く
Literal は "1mo" や "3mo" という文字列を与えますが、「this quarter」と入力した
ユーザーが 2 つ目のことを指しているとは言ってくれません。それを述べるのが spec です。
spec は引数ごとの質問、選択肢ごとの説明、関数ごとの説明、そして関数を選ぶもう 1 つの
質問を保持します。spec.json に置かれ、LLM にシグネチャから書かせることもできます。
SPEC = json.loads(Path("spec.json").read_text())
for argument in ("style", "moving_average"):
print(
json.dumps(
{argument: SPEC["functions"]["plot_price"]["arguments"][argument]}, indent=2
)
)
{
"style": {
"question": "Does the user want a plain line or candles?",
"stated": "Does the user say how the chart should be drawn, such as a line, candles, or OHLC bars?",
"options": {
"line": "a simple line through the closing prices",
"candles": "a candlestick or OHLC chart, showing each bar's open, high, low and close"
}
}
}
{
"moving_average": {
"question": "How many bars should the moving average cover - nine, twenty, or fifty?",
"stated": "Does the user ask for a moving average or a smoothed line over the candles?",
"options": {
"9": "a nine-bar moving average, a fast one",
"20": "a twenty-bar moving average",
"50": "a fifty-bar moving average, a slow one"
}
}
}
選択肢のキーは関数が受け取る文字列そのものなので、後からラベルを引数に戻す対応付けは
不要です。stated は引数を任意にします。これは、コマンドがその引数について何か言及して
いるかどうかを尋ねる 2 つ目の真偽の質問です。答えが「いいえ」なら、呼び出しはその引数を
省き、関数自身のデフォルトが適用されます。
set の引数は、メンバーごとに 1 つずつ質問を得ます。{} がメンバー名の代わりになります。
"Does the user want {} in the comparison?" はティッカーごとに 1 つの質問になります。
各質問は、ユーザーが選びそうな言葉ではなくアイデアについて書きます。一致は意味に
基づくからです。「is amd tracking nvidia lately」は、tracking も lately も
spec.json のどこにも現れないのに rolling_correlation に到達します。質問をその
パラメータの名前で名付けるのは避けます。"Which resolution?" では、コマンドが照合する
手がかりが何もありません。
spec を質問に変換する
Dispatcher は spec から質問を一度だけ組み立てます。各コマンドは、関数の選択とすべての
関数の引数を運ぶ 1 回のリクエストになり、ディスパッチャは選ばれた関数の答えだけを
読み取ります。
assistant = Dispatcher(SPEC, TOOLS, client)
print(f"{len(assistant.questions)} questions per command, among them:")
for qid in (
"__tool__",
"plot_price.style",
"plot_price.style?",
"compare_returns.symbols.NVDA",
):
question = assistant.questions[qid]
print(f" {qid:<30}{question['type']:<8}{str(question['instructions'])[:64]}")
54 questions per command, among them:
__tool__ choice What is the user asking the trading assistant to do?
plot_price.style choice Does the user want a plain line or candles?
plot_price.style? noul Does the user say how the chart should be drawn, such as a line,
compare_returns.symbols.NVDA noul Does the user want NVDA in the comparison?
14 個のコマンドを実行する
1 つのリクエストは 1 行を占め、その confidence はその呼び出しの背後で最も確信の
低い判断です。
COMMANDS = [
"show nvda 1h",
"plot rolling correlation between nvda and spy for the past month",
"when during the day does nvda trade the most",
"what moved today",
"what tickers do you have",
"how did the market do this week",
"candles for tesla with a 20 period moving average",
"compare nvda amd and msft over the past three months",
"how volatile is tsla",
"biggest losers today",
"worst drawdown for nvda this quarter, and chart it please",
"spy stats for the last month",
"show me apple daily with volume",
"is amd tracking nvidia lately",
]
CALLS = {command: assistant(command) for command in COMMANDS}
for command, call in CALLS.items():
print(f' "{command}"')
print(
f" {str(call):<66}confidence {call.confidence:.2f}"
f" tool {call.tool.probability:.2f}"
)
"show nvda 1h"
plot_price(symbol='NVDA', resolution='1h') confidence 0.78 tool 1.00
"plot rolling correlation between nvda and spy for the past month"
rolling_correlation(symbol='NVDA', benchmark='SPY', window='1mo') confidence 0.91 tool 1.00
"when during the day does nvda trade the most"
intraday_pattern(symbol='NVDA') confidence 0.53 tool 1.00
"what moved today"
top_movers(window='1d', direction='gainers') confidence 0.90 tool 0.90
"what tickers do you have"
list_symbols() confidence 1.00 tool 1.00
"how did the market do this week"
market_summary(window='1w') confidence 0.96 tool 0.99
"candles for tesla with a 20 period moving average"
plot_price(symbol='TSLA', style='candles', moving_average='20') confidence 0.69 tool 0.97
"compare nvda amd and msft over the past three months"
compare_returns(symbols=['NVDA', 'AMD', 'MSFT'], window='3mo') confidence 0.94 tool 1.00
"how volatile is tsla"
volatility(symbol='TSLA') confidence 0.96 tool 1.00
"biggest losers today"
top_movers(window='1d', direction='losers') confidence 0.98 tool 0.98
"worst drawdown for nvda this quarter, and chart it please"
drawdown(symbol='NVDA', window='3mo', plot=True) confidence 0.84 tool 0.84
"spy stats for the last month"
summary_stats(symbol='SPY', window='1mo') confidence 0.88 tool 0.88
"show me apple daily with volume"
plot_price(symbol='AAPL', resolution='1d', include_volume=True) confidence 0.75 tool 0.85
"is amd tracking nvidia lately"
rolling_correlation(symbol='AMD', benchmark='NVDA') confidence 0.82 tool 0.82
長い 2 つのコマンドは、求められたとおりに出ました。「plot rolling correlation between nvda
and spy for the past month」は 1 つの文から 4 つの引数を埋めました。そのうち symbol と
benchmark の 2 つは同じ 6 つのティッカーから取り、各ティッカーが正しい引数に収まり
ました。質問が役割を明示しているからです。測られる方、先に名付けられる と、
後に名付けられる方、基準となる物差し です。「compare nvda amd and msft over the past
three months」は 3 つのティッカーを set に入れ、残りの 3 つを外しました。
そのうち 3 つを実行します。
for command in (
"plot rolling correlation between nvda and spy for the past month",
"compare nvda amd and msft over the past three months",
"when during the day does nvda trade the most",
):
print(f'"{command}" -> {CALLS[command]}')
display(CALLS[command].run())
"plot rolling correlation between nvda and spy for the past month" -> rolling_correlation(symbol='NVDA', benchmark='SPY', window='1mo')
"compare nvda amd and msft over the past three months" -> compare_returns(symbols=['NVDA', 'AMD', 'MSFT'], window='3mo')
"when during the day does nvda trade the most" -> intraday_pattern(symbol='NVDA')
そしてテキストで答えるもの:
for command in ("how did the market do this week", "biggest losers today"):
print(f'"{command}" -> {CALLS[command]}')
print(CALLS[command].run(), "\n")
"how did the market do this week" -> market_summary(window='1w')
the board over 1w
NVDA 254.12 9.62% 389,465,563
AMD 184.20 1.51% 182,740,497
AAPL 258.71 0.97% 223,818,998
SPY 664.86 0.40% 138,617,365
MSFT 451.35 0.26% 113,427,173
TSLA 320.22 -0.97% 266,317,023
"biggest losers today" -> top_movers(window='1d', direction='losers')
top 3 losers over 1d
AMD -0.57% -> 184.20
MSFT 0.67% -> 451.35
AAPL 1.40% -> 258.71
信頼度を読む
confidence は、呼び出しの中で最も確信の低い判断を報告します。すべての積ではありません。
1 つの引数が間違っているだけで結果が台無しになるからです。積は別の質問(「すべての部分が
正しいか」)に答えるもので、関数がより多くの引数を取るほど下がります。個々の判断が
揺らいでいるかどうかとは関係ありません。
この数値がどこから来たかを、引数ごとに示します。
call = CALLS["is amd tracking nvidia lately"]
print(f'"is amd tracking nvidia lately" -> {call} confidence {call.confidence:.2f}')
for name, argument in call.arguments.items():
top = sorted(argument.distribution.items(), key=lambda kv: -kv[1])[:3]
shown = "omitted, default stands" if argument.omitted else repr(argument.value)
print(
f" {name:<12}{shown:<26}p {argument.probability:.2f} "
+ " ".join(f"{k} {v:.2f}" for k, v in top)
)
print(f" weakest argument: {call.weakest().name}")
"is amd tracking nvidia lately" -> rolling_correlation(symbol='AMD', benchmark='NVDA') confidence 0.82
symbol 'AMD' p 0.87 AMD 0.87 NVDA 0.13 AAPL 0.00
benchmark 'NVDA' p 0.78 NVDA 0.92 AMD 0.08 AAPL 0.00
window omitted, default stands p 0.96
resolution omitted, default stands p 0.99
weakest argument: benchmark
ここでは window と resolution の両方が省略されています。「lately」はどれだけ遡るか、
どの足(バー)かについて何も言っていないからです。そのため rolling_correlation は
1 か月・1 時間足という自身のデフォルトで動作します。stated の質問はそのためのものです。
これがなければ、choice は何らかのウィンドウを名指しせざるを得ず、自信を持って 1 つを
選んでいたでしょう。
playground で開く
下のリンクには、1 つのコマンドと、それが選んだ関数の質問が入っています。10 個の関数の
説明に対する choice と、rolling_correlation の 4 つの引数です。そこでコマンドを編集
すると、引数もそれに合わせて変わります。
COMMAND = "plot rolling correlation between nvda and spy for the past month"
picked = CALLS[COMMAND]
playground_link = make_playground_link(
COMMAND,
{ROUTE: assistant.questions[ROUTE]}
| {q: v for q, v in assistant.questions.items() if q.startswith(f"{picked.name}.")},
models=[TYPESAFE_MODEL],
)
display(
Markdown(
f"🔗 [Open the command and its questions in the TypeSafe playground]({playground_link})"
)
)
TypeSafe playground でコマンドとその質問を開く →