mirror of
https://github.com/tinygrad/tinygrad.git
synced 2026-08-19 19:38:27 +00:00
Compare commits
1
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
b162ab15da |
+122
@@ -0,0 +1,122 @@
|
||||
# CONTINUE.md: PAD with Invalid instead of 0
|
||||
|
||||
## Goal
|
||||
Make the low-level `Ops.PAD` pad with `Invalid` instead of `0`, while keeping
|
||||
the external `Tensor.pad` behavior unchanged.
|
||||
|
||||
## Changes made (all 3 files are modified, see `git diff`)
|
||||
|
||||
### 1. `tinygrad/schedule/indexing.py:92` — core change
|
||||
`convert_pad_to_where_to_keep_behavior_local` now uses `UOp.const(x.dtype, Invalid)`
|
||||
instead of `UOp.const(x.dtype, 0)` as the else value. This is what makes `Ops.PAD`
|
||||
pad with Invalid.
|
||||
|
||||
### 2. `tinygrad/uop/symbolic.py:87-99` — Invalid propagation rules
|
||||
Added two new rules to `pm_data_invalid` so that `where(invalid_gate, a, const_b)`
|
||||
uses `b` (the const) in don't-care positions instead of poisoning to Invalid.
|
||||
This is needed so that `_pad_constant`'s mask `where(pad(ones_bool), base, value)`
|
||||
works — the mask is a `where(valid, True, Invalid)` gate, and the else `value`
|
||||
is a const.
|
||||
|
||||
The rules are restricted to only match when the gate's valid value is a **const**
|
||||
(`UPat.cvar("x")`), to distinguish pad masks (where valid=True, a const) from
|
||||
gather masks (where valid=loaded_data, not a const). Without this restriction,
|
||||
`test_tensor_index` breaks because gather masks also create `where(cond, x, Invalid)`
|
||||
but need to keep poisoning.
|
||||
|
||||
### 3. `tinygrad/mixin/op.py:280-289` — `_pad_constant` fix
|
||||
Swapped the `value == 0` early return for `value is Invalid` early return.
|
||||
When `value is Invalid`, just return `base` (which already has Invalid from
|
||||
`Ops.PAD`). For all other values (including 0), use the mask approach:
|
||||
`where(pad(ones_bool), base, const_value)`.
|
||||
|
||||
## Current state
|
||||
- `test/unit/test_invalid_tensor.py` — **all 22 pass**
|
||||
- `test/unit/test_function.py` — **5 failures**, all multi-shard tests
|
||||
|
||||
## The remaining bug: `cat` + multi-shard
|
||||
|
||||
`cat` (op.py:716) uses `pad` + `usum` (element-wise ADD) to combine tensors:
|
||||
```python
|
||||
padded = [t.pad(...) for i,t in enumerate(tensors)]
|
||||
return padded[0].usum(*padded[1:])
|
||||
```
|
||||
|
||||
When two shards are cat'd, each is padded and then summed. The valid masks
|
||||
are **complementary** (shard 0 valid in positions 0-1, shard 1 valid in 2-3).
|
||||
|
||||
`_pad_constant` creates `where(mask_pad, data_pad, 0)` where:
|
||||
- `mask_pad = where(valid, True, Invalid)` — gate's valid value is const `True`
|
||||
- `data_pad = where(valid, data, Invalid)` — gate's valid value is loaded `data` (NOT const)
|
||||
|
||||
The new const-specific rule handles the mask pad correctly. But for the data pad,
|
||||
the gate's valid value (`data`) is not a const, so the **non-const** lift-out rule
|
||||
fires: `where(valid, where(valid, data, Invalid), 0)` → `where(valid, where(valid, data, 0), Invalid)`.
|
||||
|
||||
The `Invalid` else poisons the ADD. The binary Invalid rule lifts both gates out:
|
||||
`where(c6, data0, Invalid) + where(c8, data1, Invalid)` → `where(c6&c8, data0+data1, Invalid)`.
|
||||
|
||||
Since `c6` and `c8` are complementary, `c6&c8` is always False → result is all Invalid → 0.
|
||||
|
||||
### Master comparison
|
||||
On master, `convert_pad_to_where` uses `0` (not Invalid), so the ADD is just
|
||||
`where(c6, data0, 0) + where(c8, data1, 0)` with no Invalid, no lifting, works fine.
|
||||
|
||||
### Debug output (with changes)
|
||||
```
|
||||
c16 = c6.where(c11.index(c13), 0) # where(c6, load0, 0) — correct
|
||||
c22 = c6.where(0, c17.index(c20)) # where(c6, 0, load1) — correct
|
||||
c25 = (c6&c8).where((c16+c22), Invalid) # WRONG: c6&c8 always False → all Invalid
|
||||
```
|
||||
|
||||
### Master debug output
|
||||
```
|
||||
c13 = c6.where(c8.index(c10), 0) # where(c6, load0, 0)
|
||||
c21 = c6.where(0, c14.index(c19)) # where(c6, 0, load1)
|
||||
c22 = c13+c21 # plain ADD, no wrapper — correct
|
||||
```
|
||||
|
||||
## Suggested fix approaches
|
||||
|
||||
### Option A: General WHERE simplification rule
|
||||
Add a rule: `where(a, where(a, x, _), c)` → `where(a, x, c)`.
|
||||
When the outer and inner conditions are the same UOp, the inner else is
|
||||
unreachable. This would simplify `where(valid, where(valid, data, Invalid), 0)`
|
||||
→ `where(valid, data, 0)` before the lift-out rule can fire.
|
||||
Check if this rule already exists in `symbolic.py` — it may need to be added
|
||||
before the lift-out rules.
|
||||
|
||||
### Option B: Don't use Ops.PAD for data in `_pad_constant`
|
||||
When `value is not Invalid`, avoid creating `Ops.PAD` on the data. Use `cat`
|
||||
or `expand` to create the padded tensor directly, bypassing the Invalid
|
||||
propagation entirely.
|
||||
|
||||
### Option C: Make the lift-out rule use the outer else value
|
||||
Change the non-const lift-out rule: when `where(a, where(cond, x, Invalid), c)`
|
||||
and `c` is a const, use `c` as the else instead of `Invalid`. This is what the
|
||||
const-specific rule does, but it needs to also handle non-const gate valid values.
|
||||
|
||||
## Test commands
|
||||
```bash
|
||||
# invalid tensor tests (currently pass)
|
||||
python -m pytest test/unit/test_invalid_tensor.py -x -q -n12
|
||||
|
||||
# function tests (5 multi-shard failures)
|
||||
python -m pytest test/unit/test_function.py -x -q -n12
|
||||
|
||||
# the specific failing test
|
||||
python -m pytest test/unit/test_function.py::TestFunctionMulti::test_simple_multi_sharded -x -q
|
||||
|
||||
# debug the failing case
|
||||
DEBUG=6 python -c "
|
||||
from tinygrad import Tensor
|
||||
a = Tensor([1,2,3,4]).shard(['CPU', 'CPU:1'], axis=0)
|
||||
print(a.numpy()) # should be [1,2,3,4], gets [0,0,0,0]
|
||||
"
|
||||
```
|
||||
|
||||
## Lint/typecheck
|
||||
```bash
|
||||
python -m mypy tinygrad/
|
||||
python -m ruff check .
|
||||
```
|
||||
+26
-136
@@ -6,26 +6,6 @@ from tinygrad.helpers import partition, DEBUG, Timing, GlobalCounters, stderr_lo
|
||||
from tinygrad.viz.serve import TCPServerWithReuse, HTTPRequestHandler
|
||||
from tinygrad.llm.model import Transformer
|
||||
|
||||
def holdback(s:str, tag:str) -> int:
|
||||
# length of the suffix of s that is a prefix of tag (the tag may be split across streamed pieces)
|
||||
return max((i for i in range(1, min(len(s), len(tag))+1) if tag.startswith(s[-i:])), default=0)
|
||||
|
||||
def parse_tool_call(s:str) -> tuple[str, typing.Any]|None:
|
||||
s = s.strip()
|
||||
if s.startswith("{"): # hermes JSON format: {"name": ..., "arguments": {...}}
|
||||
try:
|
||||
call = json.loads(s)
|
||||
return call["name"], call.get("arguments", call.get("parameters", {}))
|
||||
except (json.JSONDecodeError, KeyError): return None
|
||||
# XML format: <function=name>\n<parameter=key>\nvalue\n</parameter>...</function>
|
||||
if (fm := re.match(r"<function=([^>]+)>\s*(.*?)\s*(?:</function>)?$", s, re.DOTALL)):
|
||||
args = {}
|
||||
for pm in re.finditer(r"<parameter=([^>]+)>\s*(.*?)\s*</parameter>", fm.group(2), re.DOTALL):
|
||||
try: args[pm.group(1)] = json.loads(pm.group(2))
|
||||
except json.JSONDecodeError: args[pm.group(1)] = pm.group(2)
|
||||
return fm.group(1), args
|
||||
return None
|
||||
|
||||
class SimpleTokenizer:
|
||||
def __init__(self, normal_tokens:dict[str, int], special_tokens:dict[str, int], preset:str="llama3",
|
||||
bos_id:int|None=None, eos_id:int=0, eot_id:int|None=None):
|
||||
@@ -133,75 +113,26 @@ class Handler(HTTPRequestHandler):
|
||||
def do_GET(self):
|
||||
if self.path == "/v1/models": self.send_data(json.dumps({"object":"list","data":[{"id":self.server.model_name,"object":"model"}]}).encode())
|
||||
else: self.send_data((pathlib.Path(__file__).parent / "chat.html").read_bytes(), content_type="text/html")
|
||||
def run_model(self, ids:list[int], model_name:str, include_usage=False, max_tokens:int|None=None, temperature:float=0.0,
|
||||
parse_tool_calls=False, prefill_think=False):
|
||||
def run_model(self, ids:list[int], model_name:str, include_usage=False, max_tokens:int|None=None, temperature:float=0.0):
|
||||
model, tok = self.server.model, self.server.tok
|
||||
cache_start_pos = model.get_start_pos(ids)
|
||||
stderr_log(f"{self.path} {colored('--', 'BLACK')} "
|
||||
f"in:{colored(f'{cache_start_pos:5d}', 'green')} +{len(ids)-cache_start_pos:5d} {colored('--', 'BLACK')} ")
|
||||
tmpl = {"id":f"chatcmpl-{uuid.uuid4().hex[:24]}", "object":"chat.completion.chunk", "created":int(time.time()), "model":model_name}
|
||||
def chunk(d:dict): return {"choices": [{"index":0, "delta":d, "finish_reason":None}], **tmpl}
|
||||
yield chunk({"role":"assistant", "content":""})
|
||||
yield {"choices": [{"index":0, "delta":{"role":"assistant","content":""}, "finish_reason":None}], **tmpl}
|
||||
out: list[int] = []
|
||||
finish_reason = "stop"
|
||||
st = time.perf_counter()
|
||||
dec = tok.stream_decoder()
|
||||
mode, buf, tool_text = ("reasoning" if prefill_think else "undecided"), "", ""
|
||||
def route(piece:str, final:bool=False):
|
||||
nonlocal mode, buf, tool_text
|
||||
if mode == "undecided": # decide whether the output starts with a think block
|
||||
buf += piece
|
||||
if not final and len(buf) < len("<think>") and "<think>".startswith(buf): return
|
||||
mode, piece, buf = ("reasoning", buf[len("<think>"):], "") if buf.startswith("<think>") else ("content", buf, "")
|
||||
if mode == "reasoning":
|
||||
buf += piece
|
||||
if "</think>" in buf:
|
||||
before, piece = buf.split("</think>", 1)
|
||||
if before: yield chunk({"reasoning_content":before})
|
||||
buf, mode, piece = "", "content", piece.lstrip("\n")
|
||||
else:
|
||||
hold = 0 if final else holdback(buf, "</think>")
|
||||
if (flush := buf[:len(buf)-hold]): yield chunk({"reasoning_content":flush})
|
||||
buf = buf[len(buf)-hold:]
|
||||
return
|
||||
if not parse_tool_calls:
|
||||
if piece: yield chunk({"content":piece})
|
||||
else:
|
||||
tool_text += piece
|
||||
if tool_text.startswith("<tool_call>"): return
|
||||
if "<tool_call>" in tool_text:
|
||||
before, tool_text = tool_text.split("<tool_call>", 1)
|
||||
if before: yield chunk({"content":before})
|
||||
tool_text = "<tool_call>" + tool_text
|
||||
else:
|
||||
# hold back any suffix that could be the start of a "<tool_call>" tag split across tokens
|
||||
hold = 0 if final else holdback(tool_text, "<tool_call>")
|
||||
if (flush := tool_text[:len(tool_text)-hold]): yield chunk({"content":flush})
|
||||
tool_text = tool_text[len(tool_text)-hold:]
|
||||
for next_id in model.generate(ids, temperature=temperature):
|
||||
if len(out) == 0: stderr_log(f"prefill:{(len(ids)-cache_start_pos)/((pt:=time.perf_counter())-st):4.0f} tok/s {colored('--', 'BLACK')} ")
|
||||
if tok.is_end(next_id): break
|
||||
out.append(next_id)
|
||||
yield from route(dec(next_id))
|
||||
yield {"choices": [{"index":0, "delta":{"content":dec(next_id)}, "finish_reason":None}], **tmpl}
|
||||
if max_tokens is not None and len(out) >= max_tokens:
|
||||
finish_reason = "length"
|
||||
break
|
||||
yield from route(dec(), final=True)
|
||||
if parse_tool_calls:
|
||||
tool_calls = []
|
||||
calls = [(m.group(1), m.group(0)) for m in re.finditer(r"<tool_call>\s*(.*?)\s*</tool_call>", tool_text, re.DOTALL)]
|
||||
if not calls and tool_text.startswith("<tool_call>"): calls = [(tool_text[len("<tool_call>"):], tool_text)] # unclosed tag
|
||||
for i, (inner, raw) in enumerate(calls):
|
||||
if (parsed := parse_tool_call(inner)) is None:
|
||||
stderr_log(f"failed to parse tool call: {inner[:200]}")
|
||||
yield chunk({"content":raw}) # don't silently drop output the client can't use
|
||||
else:
|
||||
name, args = parsed
|
||||
tool_calls.append({"index":i, "id":f"call_{uuid.uuid4().hex[:24]}", "type":"function",
|
||||
"function":{"name":name, "arguments":args if isinstance(args, str) else json.dumps(args)}})
|
||||
if tool_calls:
|
||||
yield chunk({"tool_calls":tool_calls})
|
||||
if finish_reason == "stop": finish_reason = "tool_calls"
|
||||
if (tail := dec()): yield {"choices": [{"index":0, "delta":{"content":tail}, "finish_reason":None}], **tmpl}
|
||||
yield {"choices": [{"index":0, "delta":{},"finish_reason":finish_reason}], **tmpl}
|
||||
if include_usage:
|
||||
yield {"choices": [], "usage": {"prompt_tokens": len(ids), "completion_tokens": len(out), "total_tokens": len(ids) + len(out)}, **tmpl}
|
||||
@@ -215,65 +146,39 @@ class Handler(HTTPRequestHandler):
|
||||
body: dict[str, typing.Any] = json.loads(raw_body.decode("utf-8"))
|
||||
if DEBUG >= 1: print(json.dumps(body, indent=2))
|
||||
if self.path == "/v1/chat/completions":
|
||||
messages, tools = body["messages"], body.get("tools")
|
||||
prefill_think = False
|
||||
if self.server.template is not None:
|
||||
# the chat template expects tool_call arguments as dicts, OpenAI clients send them as JSON strings
|
||||
norm = []
|
||||
for m in messages:
|
||||
if m.get("tool_calls"):
|
||||
m = dict(m)
|
||||
m["tool_calls"] = [{**tc, "function":{**tc["function"], "arguments":json.loads(a) if isinstance((a:=tc["function"]["arguments"]), str)
|
||||
else a}} if "function" in tc else tc for tc in m["tool_calls"]]
|
||||
norm.append(m)
|
||||
rendered = self.server.template.render(messages=norm, tools=tools, add_generation_prompt=norm[-1]["role"] != "assistant")
|
||||
prefill_think = rendered.rstrip().endswith("<think>")
|
||||
ids: list[int] = tok.encode(rendered)
|
||||
if (prefix := tok.prefix()) and ids[:len(prefix)] != prefix: ids = prefix + ids
|
||||
if norm[-1]["role"] == "assistant": # last assistant message is treated as prefill, drop its end-of-turn tokens
|
||||
end = tok.end_turn()
|
||||
if len(ids) >= len(end) and ids[-len(end):] == end: ids = ids[:-len(end)]
|
||||
else:
|
||||
if tools: stderr_log("warning: ignoring tools, install jinja2 to enable tool calling via the model's chat template")
|
||||
ids = tok.prefix()
|
||||
for i, msg in enumerate(messages):
|
||||
ids += tok.role(msg["role"])
|
||||
content = msg["content"]
|
||||
if isinstance(content, str): ids += tok.encode(content)
|
||||
elif isinstance(content, list):
|
||||
for c in content:
|
||||
if c["type"] == "text": ids += tok.encode(c["text"])
|
||||
else: raise RuntimeError(f"unhandled type: {c['type']}")
|
||||
else: raise RuntimeError(f"unknown content type: {type(content)}")
|
||||
if msg["role"] == "assistant" and i == len(messages) - 1: break
|
||||
ids += tok.end_turn()
|
||||
else: ids += tok.role("assistant")
|
||||
# extract tokens, last assistant message is treated as prefill
|
||||
ids: list[int] = tok.prefix()
|
||||
for i, msg in enumerate(body["messages"]):
|
||||
ids += tok.role(msg["role"])
|
||||
content = msg["content"]
|
||||
if isinstance(content, str): ids += tok.encode(content)
|
||||
elif isinstance(content, list):
|
||||
for c in content:
|
||||
if c["type"] == "text": ids += tok.encode(c["text"])
|
||||
else: raise RuntimeError(f"unhandled type: {c['type']}")
|
||||
else: raise RuntimeError(f"unknown content type: {type(content)}")
|
||||
if msg["role"] == "assistant" and i == len(body["messages"]) - 1: break
|
||||
ids += tok.end_turn()
|
||||
else: ids += tok.role("assistant")
|
||||
|
||||
# reply
|
||||
max_tokens = body.get("max_completion_tokens") or body.get("max_tokens")
|
||||
chunks = self.run_model(ids, body["model"], not body.get("stream") or body.get("stream_options",{}).get("include_usage", False),
|
||||
max_tokens=max_tokens, temperature=float(body.get("temperature", 0.0)), parse_tool_calls=bool(tools),
|
||||
prefill_think=prefill_think)
|
||||
max_tokens=max_tokens, temperature=float(body.get("temperature", 0.0)))
|
||||
if body.get("stream"): self.stream_json(chunks)
|
||||
else:
|
||||
out, reasoning, tool_calls, finish_reason = [], [], [], "stop"
|
||||
out, finish_reason = [], "stop"
|
||||
for c in chunks:
|
||||
if c["choices"] and (delta := c["choices"][0].get("delta", {})):
|
||||
if delta.get("content"): out.append(delta["content"])
|
||||
if delta.get("reasoning_content"): reasoning.append(delta["reasoning_content"])
|
||||
if delta.get("tool_calls"): tool_calls.extend(delta["tool_calls"])
|
||||
if c["choices"] and c["choices"][0].get("delta", {}).get("content"): out.append(c["choices"][0]["delta"]["content"])
|
||||
if c["choices"] and c["choices"][0].get("finish_reason"): finish_reason = c["choices"][0]["finish_reason"]
|
||||
message: dict[str, typing.Any] = {"role":"assistant", "content":"".join(out) or None}
|
||||
if reasoning: message["reasoning_content"] = "".join(reasoning)
|
||||
if tool_calls: message["tool_calls"] = [{k:v for k, v in tc.items() if k != "index"} for tc in tool_calls]
|
||||
self.send_data(json.dumps({**c, "object":"chat.completion",
|
||||
"choices":[{"index":0, "message":message, "finish_reason":finish_reason}]}).encode())
|
||||
"choices":[{"index":0, "message":{"role":"assistant","content":"".join(out)}, "finish_reason":finish_reason}]}).encode())
|
||||
else:
|
||||
raise RuntimeError(f"unhandled path {self.path}")
|
||||
|
||||
class LLMServer(TCPServerWithReuse):
|
||||
def __init__(self, server_address:tuple, model:Transformer, model_name:str, tok:SimpleTokenizer, template:typing.Any=None):
|
||||
self.model, self.model_name, self.tok, self.template = model, model_name, tok, template
|
||||
def __init__(self, server_address:tuple, model:Transformer, model_name:str, tok:SimpleTokenizer):
|
||||
self.model, self.model_name, self.tok = model, model_name, tok
|
||||
super().__init__(server_address, Handler)
|
||||
|
||||
def main():
|
||||
@@ -289,26 +194,11 @@ def main():
|
||||
model, kv = Transformer.from_gguf(fetch(models.get(args.model, args.model)), args.max_context)
|
||||
model_name = kv.get('general.name') or kv.get('general.basename') or args.model
|
||||
file_sizes = [y.nbytes() for y in UOp.sink(*[x.uop for x in nn.state.get_parameters(model)]).toposort() if y.op is Ops.BUFFER]
|
||||
print(f"using model \"{model_name}\" with {sum(file_sizes):,} bytes and {sum(x.numel() for x in nn.state.get_parameters(model)):,} params, "
|
||||
f"max context {args.max_context} on {nn.state.get_parameters(model)[0].device}")
|
||||
print(f"using model \"{model_name}\" with {sum(file_sizes):,} bytes and {sum(x.numel() for x in nn.state.get_parameters(model)):,} params")
|
||||
|
||||
# get tokenizer
|
||||
tok = SimpleTokenizer.from_gguf_kv(kv)
|
||||
|
||||
# compile the chat template if jinja2 is available (enables tool calling and model-specific formatting)
|
||||
template = None
|
||||
if (ct := kv.get('tokenizer.chat_template')) is not None:
|
||||
try:
|
||||
import jinja2
|
||||
env = jinja2.Environment()
|
||||
env.filters['tojson'] = lambda obj, **kwargs: json.dumps(obj) # jinja2's tojson escapes <>& for HTML safety
|
||||
env.globals['raise_exception'] = lambda msg: (_ for _ in ()).throw(RuntimeError(msg))
|
||||
env.globals['strftime_now'] = lambda fmt: time.strftime(fmt)
|
||||
for name, key in {'bos_token':'bos', 'eos_token':'eos', 'unk_token':'unknown', 'pad_token':'padding', 'sep_token':'separator'}.items():
|
||||
if (tid := kv.get(f'tokenizer.ggml.{key}_token_id')) is not None: env.globals[name] = tok.decode([tid])
|
||||
template = env.from_string(ct)
|
||||
except ImportError: stderr_log("warning: jinja2 is not installed, the model's chat template is disabled")
|
||||
|
||||
# warmup the JIT
|
||||
if args.warmup or args.serve:
|
||||
# run 2 tokens through the model twice to capture the JIT before serving
|
||||
@@ -316,7 +206,7 @@ def main():
|
||||
for _ in range(2): list(zip(range(2), model.generate([0])))
|
||||
|
||||
# start server
|
||||
if args.serve: LLMServer(('', args.serve), model, model_name, tok, template).serve_forever()
|
||||
if args.serve: LLMServer(('', args.serve), model, model_name, tok).serve_forever()
|
||||
|
||||
# do benchmark
|
||||
if args.benchmark is not None:
|
||||
|
||||
@@ -284,8 +284,8 @@ class OpMixin(ElementwiseMixin, ReduceMixin):
|
||||
X = self.shrink(tuple((-smin(pB,0),smin(pA+s,s)) for (pB,pA),s in zip(pX, self.shape))) if has_neg else self
|
||||
pads = tuple((smax(pB,0), smax(pA,0)) for pB,pA in pX) if has_neg else pX
|
||||
base = MovementMixin.pad(X, pads)
|
||||
if value == 0: return base
|
||||
if value is not Invalid: base = base.cast(least_upper_dtype(base.dtype, dtypes.from_py(value)))
|
||||
if value is Invalid: return base
|
||||
if value != 0: base = base.cast(least_upper_dtype(base.dtype, dtypes.from_py(value)))
|
||||
return MovementMixin.pad(X.const_like(1).cast(dtypes.bool), pads).where(base, base.const_like(value))
|
||||
|
||||
def _pad_circular(self, pX:tuple[tuple[sint, sint], ...]) -> Self:
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
from typing import Iterator
|
||||
import functools, itertools
|
||||
from dataclasses import dataclass, field, replace
|
||||
from tinygrad.dtype import dtypes, AddrSpace
|
||||
from tinygrad.dtype import dtypes, AddrSpace, Invalid
|
||||
from tinygrad.uop.ops import PatternMatcher, UPat, Ops, UOp, resolve, GroupOp, graph_rewrite, sint, AxisType, profile_matches
|
||||
from tinygrad.uop.ops import consumer_map_from_toposort, gate_kernel_sink
|
||||
from tinygrad.uop.symbolic import symbolic, pm_simplify_valid, pm_drop_and_clauses
|
||||
@@ -89,7 +89,7 @@ def convert_pad_to_where_to_keep_behavior_local(ctx:IndexingContext, x:UOp):
|
||||
if x not in ctx.range_map: return None
|
||||
bx = create_bufferize_and_index_based_on_ranges(ctx, x)
|
||||
valid: UOp = UOp.const(dtypes.bool, True).uprod([r.get_valid() for r in ctx.range_map[x][0]])
|
||||
return valid.where(bx.src[0], UOp.const(x.dtype, 0))
|
||||
return valid.where(bx.src[0], UOp.const(x.dtype, Invalid))
|
||||
|
||||
def convert_reduce_to_reduce_with_ranges(ctx:IndexingContext, x:UOp):
|
||||
if x.arg[1] == 0: return None
|
||||
|
||||
@@ -85,12 +85,19 @@ pm_data_invalid = PatternMatcher([
|
||||
(UPat(GroupOp.Binary, src=(UPat.var("y"), invalid_gate), name="alu"), lambda cond,x,y,alu,i: cond.where(y.alu(alu.op,x), i.cast(alu.dtype))),
|
||||
(UPat(GroupOp.Binary-GroupOp.Comparison, src=[invalid_pat, UPat()]), lambda i: i),
|
||||
# an Invalid condition poisons the whole where; a gated Invalid condition lifts the gate out
|
||||
# when the gate's valid value is a const (e.g. a pad mask: where(valid, True, Invalid)),
|
||||
# use the else value in don't-care positions so masks work
|
||||
(invalid_pat.where(UPat.var("a"), UPat()), lambda i,a: i.cast(a.dtype)),
|
||||
(UPat.var("cond").where(UPat.cvar("x"), invalid_pat).where(UPat.var("a"), UPat.cvar("b")),
|
||||
lambda cond,x,i,a,b: cond.where(x.where(a,b), b)),
|
||||
(invalid_gate.where(UPat.var("a"), UPat.var("b")), lambda cond,x,i,a,b: cond.where(x.where(a,b), i.cast(a.dtype))),
|
||||
# normalize where(cond, Invalid, val) -> where(~cond, val, Invalid)
|
||||
(UPat.var("cond").where(invalid_pat, UPat.var("val")), lambda cond, i, val: cond.logical_not().where(val, i) if val.arg != Invalid else i),
|
||||
# lift Invalid out: a.where(cond.where(x, Invalid), c) -> (~a|cond).where(a.where(x, c), Invalid)
|
||||
# when a is cond, ~a|cond is True and would drop the Invalid gate (losing the valid), so keep cond as the gate
|
||||
# when c is a const and the gate's valid value is a const (pad mask), use c in don't-care positions
|
||||
(UPat.var("a").where(UPat.var("cond").where(UPat.cvar("x"), invalid_pat), UPat.cvar("c")),
|
||||
lambda cond,i,x,a,c: (cond if a is cond else (a.logical_not()|cond)).where(a.where(x,c), c) if c.arg != Invalid else None),
|
||||
(UPat.var("a").where(invalid_gate, UPat.var("c")), lambda cond,i,x,a,c:
|
||||
(cond if a is cond else (a.logical_not()|cond)).where(a.where(x,c), i) if c.arg != Invalid else None),
|
||||
(UPat.var("a").where(UPat.var("b"), invalid_gate), lambda cond,i,x,a,b: (a|cond).where(a.where(b, x), i) if b.arg != Invalid else None),
|
||||
|
||||
Reference in New Issue
Block a user